Compare commits

..
98 Commits
Author SHA1 Message Date
欣南科技andGitHub 3f7b7da184 Merge pull request #1173 from xinnan-tech/hot-fix
update:修复立创1.6.2固件一直聆听中bug
2025-05-10 01:15:38 +08:00
hrz c0020c8769 update:修复立创1.6.2固件一直聆听中bug 2025-05-10 01:14:51 +08:00
hrzandGitHub 1c5e678cef Merge pull request #1130 from xinnan-tech/py-test
Py test
2025-05-09 18:09:31 +08:00
hrzandGitHub 3ab1b8ebc9 Merge branch 'main' into py-test 2025-05-09 18:08:48 +08:00
hrz 7635c68ccd update:增加aioconsole依赖 2025-05-09 18:01:36 +08:00
hrzandGitHub a3b562b895 Merge pull request #1113 from hsoftxl/tts-bug
tts 失败重试bug修复
2025-05-09 17:32:47 +08:00
欣南科技andGitHub 94d1493195 Merge pull request #1165 from xinnan-tech/hot-fix
fix:intent_llm的functions也需要返回成数组
2025-05-09 16:39:32 +08:00
hrz 3a5bbe32b9 fix:intent_llm的functions也需要返回成数组 2025-05-09 16:15:26 +08:00
欣南科技andGitHub 31328e092e Merge pull request #1164 from xinnan-tech/add_news
整合新的新闻插件
2025-05-09 15:26:58 +08:00
hrzandGitHub 29671f48a8 Merge branch 'main' into add_news 2025-05-09 15:25:15 +08:00
hrz 61c2782491 update:智控台,intent_llmM供应器增加functions输入框 2025-05-09 15:23:12 +08:00
hrz 15263b5223 update:更新英文版说明 2025-05-09 14:26:49 +08:00
hrz 557d1ee5b0 update:取消设备一连接就强制升级,改成绑定后再根据升级策略升级 2025-05-09 14:07:46 +08:00
hrz 777265c730 fix:参数管理保存json类型出错bug 2025-05-09 14:00:42 +08:00
hrz e76a8d547b update:使用intent_llm中的LLM进行工具内容回复 2025-05-09 13:45:01 +08:00
hrz ff98f84f18 update:intent_llm加载动态插件 2025-05-09 11:39:32 +08:00
hrzandGitHub d0b4fa5b28 Merge branch 'main' into py-test 2025-05-09 09:52:26 +08:00
hrzandGitHub 91b80f7c49 Merge pull request #1155 from xinnan-tech/web-pages-modifiy
页面优化
2025-05-09 09:51:27 +08:00
hrz 47f409246d update:intent_llm暂时不引用新闻插件,下一个版本改造intent_llm,让它能动态获取functioncall 2025-05-08 18:25:02 +08:00
hrz be1ff83297 update:通过pip来安装markitdown依赖 2025-05-08 18:23:37 +08:00
hrz 297c9e0085 update:重命名新闻插件,以平台来源来命名 2025-05-08 18:23:06 +08:00
Sakura-RanChen ec95917b4d 页面优化 2025-05-08 17:57:11 +08:00
hrzandGitHub 98bfe863fe 修复test_page.html没有hello消息,导致缺失audio_format的bug (#1151) 2025-05-08 15:08:11 +08:00
hrzandGitHub af3d00662e Merge pull request #1147 from xinnan-tech/fix-web-doubaoasr
update:更新版本号
2025-05-08 12:10:39 +08:00
欣南科技andGitHub 39803fb0df Merge pull request #1146 from xinnan-tech/fix-web-doubaoasr
修复:智控台豆包ASR缺少热词引发的bug
2025-05-08 12:09:11 +08:00
hrz 4dad5ea6c1 update:智控台添加百度ASR
修复:智控台豆包ASR缺少热词引发的bug
2025-05-08 12:07:50 +08:00
hrzandGitHub ba8cb8c6f8 Merge pull request #1125 from xinnan-tech/pcm
客户端上传编码为PCM时,服务端下发PCM格式的音频
2025-05-08 11:31:49 +08:00
hrz 64f10b28e7 update:合并main分支 2025-05-08 11:31:12 +08:00
hrzandGitHub 7f34447ace Merge branch 'main' into pcm 2025-05-08 11:18:12 +08:00
hrz c2e000f937 add:asr赋值audio_format 2025-05-08 11:11:28 +08:00
hrzandGitHub 831ab6d99b Merge pull request #1005 from JavaZeroo/add_news
feat: 添加多个新闻源,并修复大模型意图识别没有处理ActionResponse的问题
2025-05-08 10:01:33 +08:00
hrzandGitHub b246d9e567 fix:selected_module可能为空的bug (#1144)
* update:优化时间显示

* update:优化时间显示

* fix:selected_module可能为空的bug
2025-05-08 09:35:46 +08:00
hrzandGitHub 919c2ffd46 update:优化时间显示 (#1143)
* update:优化时间显示

* update:优化时间显示
2025-05-07 23:37:38 +08:00
hrzandGitHub e6d63a811e update:优化时间显示 (#1139) 2025-05-07 23:21:47 +08:00
hrzandGitHub 92227098b7 Merge pull request #1066 from xinnan-tech/manager-api-lastConnectedAtIsNull-BUG
修复了"获取用户智能体列表"中lastConnectedAt为null的bug
2025-05-07 22:57:08 +08:00
hrz 571080c1d6 update:优化最近对话时间 2025-05-07 22:56:51 +08:00
hrzandGitHub d07feb837d Merge branch 'main' into manager-api-lastConnectedAtIsNull-BUG 2025-05-07 22:45:51 +08:00
JavaZeroo 1c823c4255 feat: add news retrieval functionality and register new function 2025-05-07 21:15:06 +08:00
欣南科技andGitHub 63e378524b Merge pull request #1133 from xinnan-tech/hot-fix
update:优化日志对象
2025-05-07 18:06:57 +08:00
hrz ea5f54e421 update:优化日志对象 2025-05-07 18:06:13 +08:00
CGD ee18fbebae fix:修复“服务器运行时未处理标准输入(stdin),导致输入被缓冲,直到程序终止后才释放”的问题 2025-05-07 16:20:43 +08:00
CGD fa56d06e0c update:优化重启服务器功能 2025-05-07 16:17:50 +08:00
CGDandGitHub 72f7514114 Merge pull request #1129 from xinnan-tech/py-timeout-fix
Py timeout fix
2025-05-07 15:26:24 +08:00
Sakura-RanChen 44e1f00ffc fix:增加日志,方便意图识别 2025-05-07 15:17:34 +08:00
CGD e7054ea13f update:重启服务器功能功能 2025-05-07 14:36:08 +08:00
Sakura-RanChen 4a3ac2cfcd no message 2025-05-07 14:35:41 +08:00
Sakura-RanChen bbc31e01c4 fix:tts队列统一元组 2025-05-07 14:31:59 +08:00
玄凤科技 ba86b34a8c pcm模式,同步其他asr 2025-05-07 11:34:29 +08:00
剑雨 770a198772 获取设备最大的最近连接时间。改为在数据库里排序好后返回给系统,缓存时间修改为2分钟
--DeviceDao.java 修改方法返回值
--DeviceDao.xml 修改sql,在数据库里排序返回
--DeviceServiceImpl.java 修改缓存时间为2分钟
2025-05-07 11:04:14 +08:00
欣南科技andGitHub a26bee3696 Py update config (#1120)
* update:优化获取默认配置

* update:优化未绑定用户的连接

* update:修复智控台模式下,所选模块的日志名称

* update:优化参数配置敏感密钥的显示方式

* update:更新服务器配置并重新初始化组件
2025-05-07 09:15:52 +08:00
hrzandGitHub abb8f4f963 Merge pull request #1115 from xinnan-tech/py_tts_timeout
update: tts超时导致的文本索引混乱
2025-05-07 09:12:35 +08:00
hrz 43d2adff70 update:更新服务器配置并重新初始化组件 2025-05-07 09:11:59 +08:00
Junsen HuangandGitHub e2dffea423 Merge pull request #1104 from CaixyPromise/feature/wait-exit-fix
fix(app): 修复app.py内的wait_for_exit(),以此解决Windows环境下手动退出时,进程阻塞卡死的问题。
2025-05-07 00:48:27 +08:00
Junsen HuangandGitHub 3d5eaba46f Merge pull request #1105 from CaixyPromise/feature/mcp-exit-fix
fix(MCP): 重构MCPClient为后台协程 + AsyncExitStack管理,解决进程退出时“Attempted to exit cancel scope in a different task”错误
2025-05-07 00:47:37 +08:00
caixypromise 1bb8fbd56f resolve: merge upstream/main into feature/mcp-exit-fix and fix conflicts 2025-05-07 00:39:21 +08:00
hrz b699886953 update:优化参数配置敏感密钥的显示方式 2025-05-06 17:30:39 +08:00
hrz 12c957d48b update:修复智控台模式下,所选模块的日志名称 2025-05-06 17:14:58 +08:00
玄凤科技andGitHub 48e890c1b1 Merge pull request #1106 from kevin1sMe/feat-mcp
feat: MCP server支持使用sse模式
2025-05-06 15:38:11 +08:00
Sakura-RanChen 05331f001a update: tts超时导致的文本索引混乱 2025-05-06 15:15:27 +08:00
玄凤科技 bde260b330 PCM音频模式 2025-05-06 14:57:29 +08:00
XL 59ef51ea20 tts 失败重试bug修复 2025-05-06 14:19:05 +08:00
hrz aa77bfdfc4 update:优化未绑定用户的连接 2025-05-06 13:10:26 +08:00
hrz a2baef8911 update:优化获取默认配置 2025-05-06 13:09:43 +08:00
kevin1sMe 3fcfb65d45 update: example 2025-05-04 23:37:17 +08:00
kevin1sMe 8f266bea0d feat: MCP server支持使用sse模式 2025-05-04 23:33:07 +08:00
caixypromise 0a765f4aac fix(MCP): 重构MCPClient为后台协程 + AsyncExitStack管理,解决进程退出时的“Attempted to exit cancel scope in a different task”错误
# 变更
----
- 将所有stdio_client与ClientSession的创建/销毁都放到同一个后台 task 中
- 使用AsyncExitStack托管异步资源,cleanup时在同一task内执行exit_stack.aclose()
- 外部只通过事件通知后台task退出,避免跨协程调用cancel-scope异常
2025-05-04 23:27:47 +08:00
caixypromise ee3f0555d1 fix(app): 重写app.py内的wait_for_exit(),以此解决Windows环境下手动退出时,进程阻塞卡死的问题。
影响
----
- Ctrl‑C/kill退出时不再卡住,资源完全释放,
- Windows 与 Unix 行为一致,改动不影响正常业务逻辑。
2025-05-04 23:26:39 +08:00
hrzandGitHub 9fc1285c09 Merge pull request #1101 from xinnan-tech/hot-fix
update:删除智能体时删除聊天记录
2025-05-04 20:34:49 +08:00
欣南科技andGitHub e6e8ccee50 Merge pull request #1100 from xinnan-tech/hot-fix
Hot fix
2025-05-04 20:15:07 +08:00
hrz f9632af016 创建智能体返回智能体id 2025-05-04 20:13:32 +08:00
hrz 81359bf419 修复智控台音频播放权限bug 2025-05-04 20:03:13 +08:00
欣南科技andGitHub a144570ed7 Merge pull request #1097 from xinnan-tech/chat-history-ui
Chat history UI
2025-05-04 15:45:11 +08:00
hrz 0732b72e8f update:聊天记录音频播放 2025-05-04 15:43:28 +08:00
欣南科技andGitHub 8332445a52 Merge pull request #1096 from xinnan-tech/chat-history-ui
Chat history UI
2025-05-04 13:16:58 +08:00
hrz c9111d1bfc update:优化websocket地址验证 2025-05-04 13:15:49 +08:00
hrz e7d999278d update:优化显示 2025-05-04 13:00:56 +08:00
hrz 4043d20ea6 update:聊天记录展现 2025-05-04 12:53:36 +08:00
hrzandGitHub afe88ccfd6 Merge pull request #1091 from xinnan-tech/asr_code_clear
规范asr部分的代码,添加百度asr支持
2025-05-04 01:21:00 +08:00
hrz 56c3a04809 update:优化百度ASR文档链接 2025-05-04 01:19:54 +08:00
hrz 4530345a1f fix:设置有效的时区 2025-05-04 00:51:51 +08:00
hrzandGitHub 4a89ec8515 Merge pull request #1092 from kevin1sMe/fix-doubao-tts
fix: doubao tts token
2025-05-04 00:46:53 +08:00
欣南科技andGitHub cb5f1c6485 Merge pull request #1093 from xinnan-tech/hot-fix
add:获取聊天记录API
2025-05-03 23:49:02 +08:00
kevin1sMe 15c0677a4b fix: doubao tts token 2025-05-03 23:33:32 +08:00
hrz dc68b8148a add:获取聊天记录API 2025-05-03 23:30:39 +08:00
王华侨 1976034f12 添加百度asr支持 2025-05-03 22:21:22 +08:00
王华侨 727621fdac 规范asr部分的代码 2025-05-03 22:20:50 +08:00
欣南科技andGitHub 0f3806d358 Merge pull request #1086 from xinnan-tech/hot-fix
update:上报聊天音频
2025-05-02 16:29:48 +08:00
hrz dfb2bf3923 update:上报聊天音频 2025-05-02 16:17:18 +08:00
剑雨 b4944f23ef 添加了一个新的redis:key
--RedisKeys.java 获取设备最近最久时间的key
2025-04-30 11:08:59 +08:00
剑雨 09c17b2dc5 优化方法,把获取最近最后时间获取,放在获取设备数量前面前,可以减少一次sql查询,因为在获取时间的时候,已经顺便缓存的设备数量了
--AgentServiceImpl.java 优化:减少sql查询
2025-04-30 11:07:53 +08:00
剑雨 6f7ff9f858 优化获取这个智能体设备理的最近的最后连接时间方法,缓存时间和设备数量
--DeviceServiceImpl.java 优化方法
2025-04-30 11:05:58 +08:00
剑雨 7d74e853ae 修复智能体最近的最后连接时间为空的bug
--AgentServiceImpl.java 修复bug
2025-04-30 10:46:34 +08:00
剑雨 1d627d570d 添加获取这个智能体设备理的最近的最后连接时间定义和实现
--DeviceService.java 方法定义
--DeviceServiceImpl.java 方法实现
2025-04-30 10:45:49 +08:00
剑雨 9b0088f4f0 添加获取获取此智能体全部设备的最后连接时间,方法定义和sql,不使用mysql-plus的方法,是为了减少数据库数据传输,因为只需要最后连接时间字段
--DeviceDao.java 方法定义
--DeviceDao.xml sql
2025-04-30 10:44:54 +08:00
剑雨 f92c0313a6 添加保存用户测试方法和模拟设备连连接过来的方法,方便新开发者调试
--DeviceTest.java
2025-04-30 09:41:38 +08:00
剑雨 b9fc0da215 Merge remote-tracking branch 'origin/main' 2025-04-29 10:32:22 +08:00
剑雨 39fe057f9f Merge remote-tracking branch 'origin/main' 2025-04-17 15:56:40 +08:00
剑雨 f55d6c2e60 Merge remote-tracking branch 'origin/main' 2025-04-03 10:22:13 +08:00
106 changed files with 3703 additions and 1231 deletions
+1 -1
View File
@@ -10,7 +10,7 @@
</p> </p>
<p align="center"> <p align="center">
<a href="./README.md">English</a> <a href="./README_en.md">English</a>
· <a href="./docs/FAQ.md">常见问题</a> · <a href="./docs/FAQ.md">常见问题</a>
· <a href="https://github.com/xinnan-tech/xiaozhi-esp32-server/issues">反馈问题</a> · <a href="https://github.com/xinnan-tech/xiaozhi-esp32-server/issues">反馈问题</a>
· <a href="./README.md#%E9%83%A8%E7%BD%B2%E6%96%87%E6%A1%A3">部署文档</a> · <a href="./README.md#%E9%83%A8%E7%BD%B2%E6%96%87%E6%A1%A3">部署文档</a>
+273
View File
@@ -0,0 +1,273 @@
[![Banners](docs/images/banner1.png)](https://github.com/xinnan-tech/xiaozhi-esp32-server)
<h1 align="center">Xiaozhi Backend Service xiaozhi-esp32-server</h1>
<p align="center">
This project provides backend services for the open-source smart hardware project
<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/>
Helping you quickly set up your Xiaozhi server
</p>
<p align="center">
<a href="./README.md">中文</a>
· <a href="./docs/FAQ.md">FAQ</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="https://github.com/xinnan-tech/xiaozhi-esp32-server/releases">Release Notes</a>
</p>
<p align="center">
<a href="https://github.com/xinnan-tech/xiaozhi-esp32-server/releases">
<img alt="GitHub Contributors" src="https://img.shields.io/github/v/release/xinnan-tech/xiaozhi-esp32-server?logo=docker" />
</a>
<a href="https://github.com/xinnan-tech/xiaozhi-esp32-server/graphs/contributors">
<img alt="GitHub Contributors" src="https://img.shields.io/github/contributors/xinnan-tech/xiaozhi-esp32-server?logo=github" />
</a>
<a href="https://github.com/xinnan-tech/xiaozhi-esp32-server/issues">
<img alt="Issues" src="https://img.shields.io/github/issues/xinnan-tech/xiaozhi-esp32-server?color=0088ff" />
</a>
<a href="https://github.com/xinnan-tech/xiaozhi-esp32-server/pulls">
<img alt="GitHub pull requests" src="https://img.shields.io/github/issues-pr/xinnan-tech/xiaozhi-esp32-server?color=0088ff" />
</a>
<a href="https://github.com/xinnan-tech/xiaozhi-esp32-server/blob/main/LICENSE">
<img alt="GitHub pull requests" src="https://img.shields.io/badge/license-MIT-white?labelColor=black" />
</a>
<a href="https://github.com/xinnan-tech/xiaozhi-esp32-server">
<img alt="stars" src="https://img.shields.io/github/stars/xinnan-tech/xiaozhi-esp32-server?color=ffcb47&labelColor=black" />
</a>
</p>
---
## Target Users 👥
This project requires ESP32 hardware devices. If you have purchased ESP32-related hardware, successfully connected to Brother Xia's backend service, and want to set up your own `xiaozhi-esp32` backend service, then this project is perfect for you.
Want to see it in action? Check out these videos 🎥
<table>
<tr>
<td>
<a href="https://www.bilibili.com/video/BV1FMFyejExX" target="_blank">
<picture>
<img alt="Xiaozhi esp32 connecting to custom backend model" src="docs/images/demo1.png" />
</picture>
</a>
</td>
<td>
<a href="https://www.bilibili.com/video/BV1CDKWemEU6" target="_blank">
<picture>
<img alt="Custom voice" src="docs/images/demo2.png" />
</picture>
</a>
</td>
<td>
<a href="https://www.bilibili.com/video/BV12yA2egEaC" target="_blank">
<picture>
<img alt="Cantonese communication" src="docs/images/demo3.png" />
</picture>
</a>
</td>
<td>
<a href="https://www.bilibili.com/video/BV1pNXWYGEx1" target="_blank">
<picture>
<img alt="Home appliance control" src="docs/images/demo5.png" />
</picture>
</a>
</td>
<td>
<a href="https://www.bilibili.com/video/BV1kgA2eYEQ9" target="_blank">
<picture>
<img alt="Lowest cost configuration" src="docs/images/demo4.png" />
</picture>
</a>
</td>
</tr>
<tr>
<td>
<a href="https://www.bilibili.com/video/BV1Vy96YCE3R" target="_blank">
<picture>
<img alt="Custom voice" src="docs/images/demo6.png" />
</picture>
</a>
</td>
<td>
<a href="https://www.bilibili.com/video/BV1VC96Y5EMH" target="_blank">
<picture>
<img alt="Music playback" src="docs/images/demo7.png" />
</picture>
</a>
</td>
<td>
<a href="https://www.bilibili.com/video/BV1Z8XuYZEAS" target="_blank">
<picture>
<img alt="Weather plugin" src="docs/images/demo8.png" />
</picture>
</a>
</td>
<td>
<a href="https://www.bilibili.com/video/BV178XuYfEpi" target="_blank">
<picture>
<img alt="IOT device control" src="docs/images/demo9.png" />
</picture>
</a>
</td>
<td>
<a href="https://www.bilibili.com/video/BV17LXWYvENb" target="_blank">
<picture>
<img alt="News broadcast" src="docs/images/demo0.png" />
</picture>
</a>
</td>
</tr>
</table>
---
## Warning ⚠️
1. This project is open-source software. This software has no commercial relationship with any third-party API service providers (including but not limited to speech recognition, large models, speech synthesis, and other platforms) and does not provide any form of guarantee for their service quality or financial security.
It is recommended that users prioritize service providers with relevant business licenses and carefully read their service agreements and privacy policies. This software does not host any account keys, does not participate in fund transfers, and does not bear the risk of recharge fund losses.
2. This project's functionality is not complete and has not passed network security testing. Please do not use it in production environments. If you deploy this project for learning in a public network environment, please ensure necessary protection measures are in place.
---
## Deployment Documentation
![Banners](docs/images/banner2.png)
This project offers two deployment methods. Please choose based on your specific needs:
#### 🚀 Deployment Method Selection
| Deployment Method | Features | Use Case | Docker Deployment Guide | Source Code Deployment Guide |
|---------|------|---------|---------|---------|
| **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
Service Test Tool: https://2662r3426b.vicp.fun/test/
OTA Interface: https://2662r3426b.vicp.fun/xiaozhi/ota/
Websocket Interface: wss://2662r3426b.vicp.fun/xiaozhi/v1/
```
---
## Feature List ✨
### Implemented ✅
| Feature Module | Description |
|---------|------|
| 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 |
| 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) |
| LLM Module | Supports flexible LLM module switching, default using ChatGLMLLM, also supports Ali Bailing, DeepSeek, Ollama, etc. |
| TTS Module | Supports EdgeTTS (default), Volcano Engine Doubao TTS, and other TTS interfaces for speech synthesis |
| Memory Function | Supports ultra-long memory, local summary memory, and no memory modes for different scenarios |
| IOT Function | Supports managing registered device IOT functionality, intelligent IoT control based on dialogue context |
| Control Panel | Provides web management interface, supports agent management, user management, system configuration, etc. |
### In Development 🚧
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!
---
## Product Ecosystem 👬
Xiaozhi is an ecosystem. When using this product, you might want to check out other excellent projects in this ecosystem:
| Project Name | Project Link | 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 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 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 |
---
## Supported Platforms/Components 📋
### LLM Language Models
| Usage Method | Supported Platforms | Free Platforms |
|:---:|:---:|:---:|
| openai API | Ali Bailing, Volcano Engine Doubao, DeepSeek, ChatGLM, Gemini | ChatGLM, Gemini |
| ollama API | Ollama | - |
| dify API | Dify | - |
| fastgpt API | Fastgpt | - |
| coze API | Coze | - |
Actually, any LLM supporting openai API calls can be integrated.
---
### TTS Speech Synthesis
| 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) |
| Local Service | FishSpeech, GPT_SOVITS_V2, GPT_SOVITS_V3, MinimaxTTS | FishSpeech, GPT_SOVITS_V2, GPT_SOVITS_V3, MinimaxTTS |
---
### VAD Voice Activity Detection
| Type | Platform Name | Usage Method | Pricing | Notes |
|:---:|:---------:|:----:|:----:|:--:|
| VAD | SileroVAD | Local Use | Free | |
---
### ASR Speech Recognition
| Usage Method | Supported Platforms | Free Platforms |
|:---:|:---:|:---:|
| Local Use | FunASR, SherpaASR | FunASR, SherpaASR |
| API Calls | DoubaoASR, FunASRServer, TencentASR, AliyunASR | FunASRServer |
---
### Memory Storage
| Type | Platform Name | Usage Method | Pricing | Notes |
|:------:|:---------------:|:----:|:---------:|:--:|
| Memory | mem0ai | API Calls | 1000 calls/month quota | |
| Memory | mem_local_short | Local Summary | Free | |
---
### Intent Recognition
| Type | Platform Name | Usage Method | Pricing | Notes |
|:------:|:-------------:|:----:|:-------:|:---------------------:|
| Intent | intent_llm | API Calls | 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 |
---
## Acknowledgments 🙏
| 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_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_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_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_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 |
<a href="https://star-history.com/#xinnan-tech/xiaozhi-esp32-server&Date">
<picture>
<source media="(prefers-color-scheme: dark)" srcset="https://api.star-history.com/svg?repos=xinnan-tech/xiaozhi-esp32-server&type=Date&theme=dark" />
<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" />
</picture>
</a>
+1 -1
View File
@@ -213,7 +213,7 @@ CREATE DATABASE xiaozhi_esp32_server CHARACTER SET utf8mb4 COLLATE utf8mb4_unico
如果还没有MySQL,你可以通过docker安装mysql 如果还没有MySQL,你可以通过docker安装mysql
``` ```
docker run --name xiaozhi-esp32-server-db -e MYSQL_ROOT_PASSWORD=123456 -p 3306:3306 -e MYSQL_DATABASE=xiaozhi_esp32_server -e MYSQL_INITDB_ARGS="--character-set-server=utf8mb4 --collation-server=utf8mb4_unicode_ci" -d mysql:latest docker run --name xiaozhi-esp32-server-db -e MYSQL_ROOT_PASSWORD=123456 -p 3306:3306 -e MYSQL_DATABASE=xiaozhi_esp32_server -e MYSQL_INITDB_ARGS="--character-set-server=utf8mb4 --collation-server=utf8mb4_unicode_ci" -e TZ=Asia/Shanghai -d mysql:latest
``` ```
## 2.安装redis ## 2.安装redis
+1 -1
View File
@@ -16,4 +16,4 @@ https://ccnphfhqs21z.feishu.cn/wiki/M0XiwldO9iJwHikpXD5cEx71nKh
# manager-web 、manager-api接口协议 # manager-web 、manager-api接口协议
https://2662r3426b.vicp.fun/xiaozhi/v1/doc.html https://2662r3426b.vicp.fun/xiaozhi/doc.html
+4
View File
@@ -197,6 +197,10 @@
<artifactId>lombok</artifactId> <artifactId>lombok</artifactId>
<optional>true</optional> <optional>true</optional>
</dependency> </dependency>
<dependency>
<groupId>com.fasterxml.jackson.datatype</groupId>
<artifactId>jackson-datatype-jsr310</artifactId>
</dependency>
</dependencies> </dependencies>
<!-- 阿里云maven仓库 --> <!-- 阿里云maven仓库 -->
@@ -177,5 +177,5 @@ public interface Constant {
/** /**
* 版本号 * 版本号
*/ */
public static final String VERSION = "0.3.13"; public static final String VERSION = "0.4.2";
} }
@@ -62,6 +62,13 @@ public class RedisKeys {
return "agent:device:count:" + id; return "agent:device:count:" + id;
} }
/**
* 获取智能体最后连接时间缓存key
*/
public static String getAgentDeviceLastConnectedAtById(String id) {
return "agent:device:lastConnected:" + id;
}
/** /**
* 获取系统配置缓存key * 获取系统配置缓存key
*/ */
@@ -103,4 +110,11 @@ public class RedisKeys {
public static String getDictDataByTypeKey(String dictType) { public static String getDictDataByTypeKey(String dictType) {
return "sys:dict:data:" + dictType; return "sys:dict:data:" + dictType;
} }
/**
* 获取智能体音频ID的缓存key
*/
public static String getAgentAudioIdKey(String uuid) {
return "agent:audio:id:" + uuid;
}
} }
@@ -3,8 +3,13 @@ package xiaozhi.modules.agent.controller;
import java.util.Date; import java.util.Date;
import java.util.List; import java.util.List;
import java.util.Map; import java.util.Map;
import java.util.UUID;
import org.apache.commons.lang3.StringUtils;
import org.apache.shiro.authz.annotation.RequiresPermissions; import org.apache.shiro.authz.annotation.RequiresPermissions;
import org.springframework.http.HttpHeaders;
import org.springframework.http.MediaType;
import org.springframework.http.ResponseEntity;
import org.springframework.web.bind.annotation.DeleteMapping; import org.springframework.web.bind.annotation.DeleteMapping;
import org.springframework.web.bind.annotation.GetMapping; import org.springframework.web.bind.annotation.GetMapping;
import org.springframework.web.bind.annotation.PathVariable; import org.springframework.web.bind.annotation.PathVariable;
@@ -25,14 +30,20 @@ import jakarta.validation.Valid;
import lombok.AllArgsConstructor; import lombok.AllArgsConstructor;
import xiaozhi.common.constant.Constant; import xiaozhi.common.constant.Constant;
import xiaozhi.common.page.PageData; import xiaozhi.common.page.PageData;
import xiaozhi.common.redis.RedisKeys;
import xiaozhi.common.redis.RedisUtils;
import xiaozhi.common.user.UserDetail; import xiaozhi.common.user.UserDetail;
import xiaozhi.common.utils.ConvertUtils; import xiaozhi.common.utils.ConvertUtils;
import xiaozhi.common.utils.Result; import xiaozhi.common.utils.Result;
import xiaozhi.modules.agent.dto.AgentChatHistoryDTO;
import xiaozhi.modules.agent.dto.AgentChatSessionDTO;
import xiaozhi.modules.agent.dto.AgentCreateDTO; import xiaozhi.modules.agent.dto.AgentCreateDTO;
import xiaozhi.modules.agent.dto.AgentDTO; import xiaozhi.modules.agent.dto.AgentDTO;
import xiaozhi.modules.agent.dto.AgentUpdateDTO; import xiaozhi.modules.agent.dto.AgentUpdateDTO;
import xiaozhi.modules.agent.entity.AgentEntity; import xiaozhi.modules.agent.entity.AgentEntity;
import xiaozhi.modules.agent.entity.AgentTemplateEntity; import xiaozhi.modules.agent.entity.AgentTemplateEntity;
import xiaozhi.modules.agent.service.AgentChatAudioService;
import xiaozhi.modules.agent.service.AgentChatHistoryService;
import xiaozhi.modules.agent.service.AgentService; import xiaozhi.modules.agent.service.AgentService;
import xiaozhi.modules.agent.service.AgentTemplateService; import xiaozhi.modules.agent.service.AgentTemplateService;
import xiaozhi.modules.device.service.DeviceService; import xiaozhi.modules.device.service.DeviceService;
@@ -46,6 +57,9 @@ public class AgentController {
private final AgentService agentService; private final AgentService agentService;
private final AgentTemplateService agentTemplateService; private final AgentTemplateService agentTemplateService;
private final DeviceService deviceService; private final DeviceService deviceService;
private final AgentChatHistoryService agentChatHistoryService;
private final AgentChatAudioService agentChatAudioService;
private final RedisUtils redisUtils;
@GetMapping("/list") @GetMapping("/list")
@Operation(summary = "获取用户智能体列表") @Operation(summary = "获取用户智能体列表")
@@ -80,7 +94,7 @@ public class AgentController {
@PostMapping @PostMapping
@Operation(summary = "创建智能体") @Operation(summary = "创建智能体")
@RequiresPermissions("sys:role:normal") @RequiresPermissions("sys:role:normal")
public Result<Void> save(@RequestBody @Valid AgentCreateDTO dto) { public Result<String> save(@RequestBody @Valid AgentCreateDTO dto) {
AgentEntity entity = ConvertUtils.sourceToTarget(dto, AgentEntity.class); AgentEntity entity = ConvertUtils.sourceToTarget(dto, AgentEntity.class);
// 获取默认模板 // 获取默认模板
@@ -108,7 +122,7 @@ public class AgentController {
// ID、智能体编码和排序会在Service层自动生成 // ID、智能体编码和排序会在Service层自动生成
agentService.insert(entity); agentService.insert(entity);
return new Result<>(); return new Result<String>().ok(entity.getId());
} }
@PutMapping("/{id}") @PutMapping("/{id}")
@@ -178,6 +192,8 @@ public class AgentController {
public Result<Void> delete(@PathVariable String id) { public Result<Void> delete(@PathVariable String id) {
// 先删除关联的设备 // 先删除关联的设备
deviceService.deleteByAgentId(id); deviceService.deleteByAgentId(id);
// 删除关联的聊天记录
agentChatHistoryService.deleteByAgentId(id);
// 再删除智能体 // 再删除智能体
agentService.deleteById(id); agentService.deleteById(id);
return new Result<>(); return new Result<>();
@@ -192,4 +208,71 @@ public class AgentController {
return new Result<List<AgentTemplateEntity>>().ok(list); return new Result<List<AgentTemplateEntity>>().ok(list);
} }
@GetMapping("/{id}/sessions")
@Operation(summary = "获取智能体会话列表")
@RequiresPermissions("sys:role:normal")
@Parameters({
@Parameter(name = Constant.PAGE, description = "当前页码,从1开始", required = true),
@Parameter(name = Constant.LIMIT, description = "每页显示记录数", required = true),
})
public Result<PageData<AgentChatSessionDTO>> getAgentSessions(
@PathVariable("id") String id,
@Parameter(hidden = true) @RequestParam Map<String, Object> params) {
params.put("agentId", id);
PageData<AgentChatSessionDTO> page = agentChatHistoryService.getSessionListByAgentId(params);
return new Result<PageData<AgentChatSessionDTO>>().ok(page);
}
@GetMapping("/{id}/chat-history/{sessionId}")
@Operation(summary = "获取智能体聊天记录")
@RequiresPermissions("sys:role:normal")
public Result<List<AgentChatHistoryDTO>> getAgentChatHistory(
@PathVariable("id") String id,
@PathVariable("sessionId") String sessionId) {
// 获取当前用户
UserDetail user = SecurityUser.getUser();
// 检查权限
if (!agentService.checkAgentPermission(id, user.getId())) {
return new Result<List<AgentChatHistoryDTO>>().error("没有权限查看该智能体的聊天记录");
}
// 查询聊天记录
List<AgentChatHistoryDTO> result = agentChatHistoryService.getChatHistoryBySessionId(id, sessionId);
return new Result<List<AgentChatHistoryDTO>>().ok(result);
}
@PostMapping("/audio/{audioId}")
@Operation(summary = "获取音频下载ID")
@RequiresPermissions("sys:role:normal")
public Result<String> getAudioId(@PathVariable("audioId") String audioId) {
byte[] audioData = agentChatAudioService.getAudio(audioId);
if (audioData == null) {
return new Result<String>().error("音频不存在");
}
String uuid = UUID.randomUUID().toString();
redisUtils.set(RedisKeys.getAgentAudioIdKey(uuid), audioId);
return new Result<String>().ok(uuid);
}
@GetMapping("/play/{uuid}")
@Operation(summary = "播放音频")
public ResponseEntity<byte[]> playAudio(@PathVariable("uuid") String uuid) {
String audioId = (String) redisUtils.get(RedisKeys.getAgentAudioIdKey(uuid));
if (StringUtils.isBlank(audioId)) {
return ResponseEntity.notFound().build();
}
byte[] audioData = agentChatAudioService.getAudio(audioId);
if (audioData == null) {
return ResponseEntity.notFound().build();
}
redisUtils.delete(RedisKeys.getAgentAudioIdKey(uuid));
return ResponseEntity.ok()
.contentType(MediaType.APPLICATION_OCTET_STREAM)
.header(HttpHeaders.CONTENT_DISPOSITION, "attachment; filename=\"play.wav\"")
.body(audioData);
}
} }
@@ -1,7 +1,9 @@
package xiaozhi.modules.agent.dao; package xiaozhi.modules.agent.dao;
import com.baomidou.mybatisplus.core.mapper.BaseMapper;
import org.apache.ibatis.annotations.Mapper; import org.apache.ibatis.annotations.Mapper;
import com.baomidou.mybatisplus.core.mapper.BaseMapper;
import xiaozhi.modules.agent.entity.AgentChatHistoryEntity; import xiaozhi.modules.agent.entity.AgentChatHistoryEntity;
/** /**
@@ -13,4 +15,17 @@ import xiaozhi.modules.agent.entity.AgentChatHistoryEntity;
*/ */
@Mapper @Mapper
public interface AiAgentChatHistoryDao extends BaseMapper<AgentChatHistoryEntity> { public interface AiAgentChatHistoryDao extends BaseMapper<AgentChatHistoryEntity> {
/**
* 根据智能体ID删除音频
*
* @param agentId 智能体ID
*/
void deleteAudioByAgentId(String agentId);
/**
* 根据智能体ID删除聊天历史记录
*
* @param agentId 智能体ID
*/
void deleteHistoryByAgentId(String agentId);
} }
@@ -0,0 +1,28 @@
package xiaozhi.modules.agent.dto;
import java.util.Date;
import io.swagger.v3.oas.annotations.media.Schema;
import lombok.Data;
/**
* 智能体聊天记录DTO
*/
@Data
@Schema(description = "智能体聊天记录")
public class AgentChatHistoryDTO {
@Schema(description = "创建时间")
private Date createdAt;
@Schema(description = "消息类型: 1-用户, 2-智能体")
private Byte chatType;
@Schema(description = "聊天内容")
private String content;
@Schema(description = "音频ID")
private String audioId;
@Schema(description = "MAC地址")
private String macAddress;
}
@@ -26,6 +26,6 @@ public class AgentChatHistoryReportDTO {
@Schema(description = "聊天内容", example = "你好呀") @Schema(description = "聊天内容", example = "你好呀")
@NotBlank @NotBlank
private String content; private String content;
@Schema(description = "文件数据(opus编码)", example = "") @Schema(description = "base64编码的opus音频数据", example = "")
private String opusDataBase64; private String audioBase64;
} }
@@ -0,0 +1,26 @@
package xiaozhi.modules.agent.dto;
import java.time.LocalDateTime;
import lombok.Data;
/**
* 智能体会话列表DTO
*/
@Data
public class AgentChatSessionDTO {
/**
* 会话ID
*/
private String sessionId;
/**
* 会话时间
*/
private LocalDateTime createdAt;
/**
* 聊天条数
*/
private Integer chatCount;
}
@@ -67,12 +67,6 @@ public class AgentChatHistoryEntity {
@TableField(value = "audio_id") @TableField(value = "audio_id")
private String audioId; private String audioId;
/**
* 音频URL
*/
@TableField(value = "audio_url")
private String audioUrl;
/** /**
* 创建时间 * 创建时间
*/ */
@@ -19,4 +19,12 @@ public interface AgentChatAudioService extends IService<AgentChatAudioEntity> {
* @return 音频ID * @return 音频ID
*/ */
String saveAudio(byte[] audioData); String saveAudio(byte[] audioData);
/**
* 获取音频数据
*
* @param audioId 音频ID
* @return 音频数据
*/
byte[] getAudio(String audioId);
} }
@@ -1,6 +1,13 @@
package xiaozhi.modules.agent.service; package xiaozhi.modules.agent.service;
import java.util.List;
import java.util.Map;
import com.baomidou.mybatisplus.extension.service.IService; import com.baomidou.mybatisplus.extension.service.IService;
import xiaozhi.common.page.PageData;
import xiaozhi.modules.agent.dto.AgentChatHistoryDTO;
import xiaozhi.modules.agent.dto.AgentChatSessionDTO;
import xiaozhi.modules.agent.entity.AgentChatHistoryEntity; import xiaozhi.modules.agent.entity.AgentChatHistoryEntity;
/** /**
@@ -11,4 +18,28 @@ import xiaozhi.modules.agent.entity.AgentChatHistoryEntity;
* @since 1.0.0 * @since 1.0.0
*/ */
public interface AgentChatHistoryService extends IService<AgentChatHistoryEntity> { public interface AgentChatHistoryService extends IService<AgentChatHistoryEntity> {
/**
* 根据智能体ID获取会话列表
*
* @param params 查询参数,包含agentId、page、limit
* @return 分页的会话列表
*/
PageData<AgentChatSessionDTO> getSessionListByAgentId(Map<String, Object> params);
/**
* 根据会话ID获取聊天记录列表
*
* @param agentId 智能体ID
* @param sessionId 会话ID
* @return 聊天记录列表
*/
List<AgentChatHistoryDTO> getChatHistoryBySessionId(String agentId, String sessionId);
/**
* 根据智能体ID删除聊天记录
*
* @param agentId 智能体ID
*/
void deleteByAgentId(String agentId);
} }
@@ -8,36 +8,56 @@ import xiaozhi.common.service.BaseService;
import xiaozhi.modules.agent.dto.AgentDTO; import xiaozhi.modules.agent.dto.AgentDTO;
import xiaozhi.modules.agent.entity.AgentEntity; import xiaozhi.modules.agent.entity.AgentEntity;
/**
* 智能体表处理service
*
* @author Goody
* @version 1.0, 2025/4/30
* @since 1.0.0
*/
public interface AgentService extends BaseService<AgentEntity> { public interface AgentService extends BaseService<AgentEntity> {
/** /**
* 管理员获取所有智能体列表(分页) * 获取管理员智能体列表
*
* @param params 查询参数
* @return 分页数据
*/ */
PageData<AgentEntity> adminAgentList(Map<String, Object> params); PageData<AgentEntity> adminAgentList(Map<String, Object> params);
/** /**
* 获取智能体详情 * 根据ID获取智能体
*
* @param id 智能体ID
* @return 智能体实体
*/ */
AgentEntity getAgentById(String id); AgentEntity getAgentById(String id);
/** /**
* 删除这个用户的所有 * 插入智能体
* *
* @param userId * @param entity 智能体实体
* @return 是否成功
*/
boolean insert(AgentEntity entity);
/**
* 根据用户ID删除智能体
*
* @param userId 用户ID
*/ */
void deleteAgentByUserId(Long userId); void deleteAgentByUserId(Long userId);
/** /**
* 获取用户智能体列表 * 获取用户智能体列表
* *
* @param userId * @param userId 用户ID
* @return * @return 智能体列表
*/ */
List<AgentDTO> getUserAgents(Long userId); List<AgentDTO> getUserAgents(Long userId);
/** /**
* 获取智能体设备数量 * 根据智能体ID获取设备数量
* *
* @param agentId 智能体ID * @param agentId 智能体ID
* @return 设备数量 * @return 设备数量
*/ */
@@ -50,4 +70,13 @@ public interface AgentService extends BaseService<AgentEntity> {
* @return 默认智能体信息,不存在时返回null * @return 默认智能体信息,不存在时返回null
*/ */
AgentEntity getDefaultAgentByMacAddress(String macAddress); AgentEntity getDefaultAgentByMacAddress(String macAddress);
/**
* 检查用户是否有权限访问智能体
*
* @param agentId 智能体ID
* @param userId 用户ID
* @return 是否有权限
*/
boolean checkAgentPermission(String agentId, Long userId);
} }
@@ -1,10 +1,15 @@
package xiaozhi.modules.agent.service.biz.impl; package xiaozhi.modules.agent.service.biz.impl;
import java.util.Base64;
import java.util.Date;
import org.springframework.stereotype.Service; import org.springframework.stereotype.Service;
import org.springframework.transaction.annotation.Transactional; import org.springframework.transaction.annotation.Transactional;
import lombok.RequiredArgsConstructor; import lombok.RequiredArgsConstructor;
import lombok.extern.slf4j.Slf4j; import lombok.extern.slf4j.Slf4j;
import xiaozhi.common.redis.RedisKeys;
import xiaozhi.common.redis.RedisUtils;
import xiaozhi.modules.agent.dto.AgentChatHistoryReportDTO; import xiaozhi.modules.agent.dto.AgentChatHistoryReportDTO;
import xiaozhi.modules.agent.entity.AgentChatHistoryEntity; import xiaozhi.modules.agent.entity.AgentChatHistoryEntity;
import xiaozhi.modules.agent.entity.AgentEntity; import xiaozhi.modules.agent.entity.AgentEntity;
@@ -27,6 +32,7 @@ public class AgentChatHistoryBizServiceImpl implements AgentChatHistoryBizServic
private final AgentService agentService; private final AgentService agentService;
private final AgentChatHistoryService agentChatHistoryService; private final AgentChatHistoryService agentChatHistoryService;
private final AgentChatAudioService agentChatAudioService; private final AgentChatAudioService agentChatAudioService;
private final RedisUtils redisUtils;
/** /**
* 处理聊天记录上报,包括文件上传和相关信息记录 * 处理聊天记录上报,包括文件上传和相关信息记录
@@ -43,12 +49,11 @@ public class AgentChatHistoryBizServiceImpl implements AgentChatHistoryBizServic
// 1. base64解码report.getOpusDataBase64(),存入ai_agent_chat_audio表 // 1. base64解码report.getOpusDataBase64(),存入ai_agent_chat_audio表
String audioId = null; String audioId = null;
if (report.getOpusDataBase64() != null && !report.getOpusDataBase64().isEmpty()) { if (report.getAudioBase64() != null && !report.getAudioBase64().isEmpty()) {
try { try {
// TODO: 需要考虑保留什么格式的音频数据,比如是opus还是wave byte[] audioData = Base64.getDecoder().decode(report.getAudioBase64());
// byte[] audioData = Base64.getDecoder().decode(report.getOpusDataBase64()); audioId = agentChatAudioService.saveAudio(audioData);
// audioId = agentChatAudioService.saveAudio(audioData); log.info("音频数据保存成功,audioId={}", audioId);
// log.info("音频数据保存成功,audioId={}", audioId);
} catch (Exception e) { } catch (Exception e) {
log.error("音频数据保存失败", e); log.error("音频数据保存失败", e);
return false; return false;
@@ -76,6 +81,8 @@ public class AgentChatHistoryBizServiceImpl implements AgentChatHistoryBizServic
// 3. 保存数据 // 3. 保存数据
agentChatHistoryService.save(entity); agentChatHistoryService.save(entity);
// 4. 更新设备最后对话时间
redisUtils.set(RedisKeys.getAgentDeviceLastConnectedAtById(agentId), new Date());
return Boolean.TRUE; return Boolean.TRUE;
} }
} }
@@ -25,4 +25,10 @@ public class AgentChatAudioServiceImpl extends ServiceImpl<AiAgentChatAudioDao,
save(entity); save(entity);
return entity.getId(); return entity.getId();
} }
@Override
public byte[] getAudio(String audioId) {
AgentChatAudioEntity entity = getById(audioId);
return entity != null ? entity.getAudio() : null;
}
} }
@@ -1,19 +1,85 @@
package xiaozhi.modules.agent.service.impl; package xiaozhi.modules.agent.service.impl;
import com.baomidou.mybatisplus.extension.service.impl.ServiceImpl; import java.time.LocalDateTime;
import org.springframework.stereotype.Service; import java.util.List;
import xiaozhi.modules.agent.dao.AiAgentChatHistoryDao; import java.util.Map;
import xiaozhi.modules.agent.entity.AgentChatHistoryEntity; import java.util.stream.Collectors;
import xiaozhi.modules.agent.service.AgentChatHistoryService;
import org.springframework.stereotype.Service;
/** import org.springframework.transaction.annotation.Transactional;
* 智能体聊天记录表处理service {@link AgentChatHistoryService} impl
* import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper;
* @author Goody import com.baomidou.mybatisplus.core.metadata.IPage;
* @version 1.0, 2025/4/30 import com.baomidou.mybatisplus.extension.plugins.pagination.Page;
* @since 1.0.0 import com.baomidou.mybatisplus.extension.service.impl.ServiceImpl;
*/
@Service import xiaozhi.common.constant.Constant;
public class AgentChatHistoryServiceImpl extends ServiceImpl<AiAgentChatHistoryDao, AgentChatHistoryEntity> implements AgentChatHistoryService { import xiaozhi.common.page.PageData;
import xiaozhi.common.utils.ConvertUtils;
} import xiaozhi.modules.agent.dao.AiAgentChatHistoryDao;
import xiaozhi.modules.agent.dto.AgentChatHistoryDTO;
import xiaozhi.modules.agent.dto.AgentChatSessionDTO;
import xiaozhi.modules.agent.entity.AgentChatHistoryEntity;
import xiaozhi.modules.agent.service.AgentChatHistoryService;
/**
* 智能体聊天记录表处理service {@link AgentChatHistoryService} impl
*
* @author Goody
* @version 1.0, 2025/4/30
* @since 1.0.0
*/
@Service
public class AgentChatHistoryServiceImpl extends ServiceImpl<AiAgentChatHistoryDao, AgentChatHistoryEntity>
implements AgentChatHistoryService {
@Override
public PageData<AgentChatSessionDTO> getSessionListByAgentId(Map<String, Object> params) {
String agentId = (String) params.get("agentId");
int page = Integer.parseInt(params.get(Constant.PAGE).toString());
int limit = Integer.parseInt(params.get(Constant.LIMIT).toString());
// 构建查询条件
QueryWrapper<AgentChatHistoryEntity> wrapper = new QueryWrapper<>();
wrapper.select("session_id", "MAX(created_at) as created_at", "COUNT(*) as chat_count")
.eq("agent_id", agentId)
.groupBy("session_id")
.orderByDesc("created_at");
// 执行分页查询
Page<Map<String, Object>> pageParam = new Page<>(page, limit);
IPage<Map<String, Object>> result = this.baseMapper.selectMapsPage(pageParam, wrapper);
List<AgentChatSessionDTO> records = result.getRecords().stream().map(map -> {
AgentChatSessionDTO dto = new AgentChatSessionDTO();
dto.setSessionId((String) map.get("session_id"));
dto.setCreatedAt((LocalDateTime) map.get("created_at"));
dto.setChatCount(((Number) map.get("chat_count")).intValue());
return dto;
}).collect(Collectors.toList());
return new PageData<>(records, result.getTotal());
}
@Override
public List<AgentChatHistoryDTO> getChatHistoryBySessionId(String agentId, String sessionId) {
// 构建查询条件
QueryWrapper<AgentChatHistoryEntity> wrapper = new QueryWrapper<>();
wrapper.eq("agent_id", agentId)
.eq("session_id", sessionId)
.orderByAsc("created_at");
// 查询聊天记录
List<AgentChatHistoryEntity> historyList = list(wrapper);
// 转换为DTO
return ConvertUtils.sourceToTarget(historyList, AgentChatHistoryDTO.class);
}
@Override
@Transactional(rollbackFor = Exception.class)
public void deleteByAgentId(String agentId) {
baseMapper.deleteAudioByAgentId(agentId);
baseMapper.deleteHistoryByAgentId(agentId);
}
}
@@ -6,13 +6,13 @@ import java.util.UUID;
import java.util.stream.Collectors; import java.util.stream.Collectors;
import org.apache.commons.lang3.StringUtils; import org.apache.commons.lang3.StringUtils;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.stereotype.Service; import org.springframework.stereotype.Service;
import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper; import com.baomidou.mybatisplus.core.conditions.query.QueryWrapper;
import com.baomidou.mybatisplus.core.conditions.update.UpdateWrapper; import com.baomidou.mybatisplus.core.conditions.update.UpdateWrapper;
import com.baomidou.mybatisplus.core.metadata.IPage; import com.baomidou.mybatisplus.core.metadata.IPage;
import lombok.AllArgsConstructor;
import xiaozhi.common.page.PageData; import xiaozhi.common.page.PageData;
import xiaozhi.common.redis.RedisKeys; import xiaozhi.common.redis.RedisKeys;
import xiaozhi.common.redis.RedisUtils; import xiaozhi.common.redis.RedisUtils;
@@ -21,25 +21,20 @@ import xiaozhi.modules.agent.dao.AgentDao;
import xiaozhi.modules.agent.dto.AgentDTO; import xiaozhi.modules.agent.dto.AgentDTO;
import xiaozhi.modules.agent.entity.AgentEntity; import xiaozhi.modules.agent.entity.AgentEntity;
import xiaozhi.modules.agent.service.AgentService; import xiaozhi.modules.agent.service.AgentService;
import xiaozhi.modules.device.service.DeviceService;
import xiaozhi.modules.model.service.ModelConfigService; import xiaozhi.modules.model.service.ModelConfigService;
import xiaozhi.modules.security.user.SecurityUser;
import xiaozhi.modules.sys.enums.SuperAdminEnum;
import xiaozhi.modules.timbre.service.TimbreService; import xiaozhi.modules.timbre.service.TimbreService;
@Service @Service
@AllArgsConstructor
public class AgentServiceImpl extends BaseServiceImpl<AgentDao, AgentEntity> implements AgentService { public class AgentServiceImpl extends BaseServiceImpl<AgentDao, AgentEntity> implements AgentService {
private final AgentDao agentDao; private final AgentDao agentDao;
private final TimbreService timbreModelService;
@Autowired private final ModelConfigService modelConfigService;
private TimbreService timbreModelService; private final RedisUtils redisUtils;
private final DeviceService deviceService;
@Autowired
private ModelConfigService modelConfigService;
@Autowired
private RedisUtils redisUtils;
public AgentServiceImpl(AgentDao agentDao) {
this.agentDao = agentDao;
}
@Override @Override
public PageData<AgentEntity> adminAgentList(Map<String, Object> params) { public PageData<AgentEntity> adminAgentList(Map<String, Object> params) {
@@ -101,9 +96,11 @@ public class AgentServiceImpl extends BaseServiceImpl<AgentDao, AgentEntity> imp
// 获取 TTS 音色名称 // 获取 TTS 音色名称
dto.setTtsVoiceName(timbreModelService.getTimbreNameById(agent.getTtsVoiceId())); dto.setTtsVoiceName(timbreModelService.getTimbreNameById(agent.getTtsVoiceId()));
// 获取智能体最近的最后连接时长
dto.setLastConnectedAt(deviceService.getLatestLastConnectionTime(agent.getId()));
// 获取设备数量 // 获取设备数量
dto.setDeviceCount(getDeviceCountByAgentId(agent.getId())); dto.setDeviceCount(getDeviceCountByAgentId(agent.getId()));
return dto; return dto;
}).collect(Collectors.toList()); }).collect(Collectors.toList());
} }
@@ -138,4 +135,21 @@ public class AgentServiceImpl extends BaseServiceImpl<AgentDao, AgentEntity> imp
} }
return agentDao.getDefaultAgentByMacAddress(macAddress); return agentDao.getDefaultAgentByMacAddress(macAddress);
} }
@Override
public boolean checkAgentPermission(String agentId, Long userId) {
// 获取智能体信息
AgentEntity agent = getAgentById(agentId);
if (agent == null) {
return false;
}
// 如果是超级管理员,直接返回true
if (SecurityUser.getUser().getSuperAdmin() == SuperAdminEnum.YES.value()) {
return true;
}
// 检查是否是智能体的所有者
return userId.equals(agent.getUserId());
}
} }
@@ -61,15 +61,15 @@ public class ConfigServiceImpl implements ConfigService {
// 构建模块配置 // 构建模块配置
buildModuleConfig( buildModuleConfig(
agent.getAgentName(), null,
null, null,
null, null,
agent.getVadModelId(), agent.getVadModelId(),
agent.getAsrModelId(), agent.getAsrModelId(),
agent.getLlmModelId(), null,
agent.getTtsModelId(), null,
agent.getMemModelId(), null,
agent.getIntentModelId(), null,
result, result,
isCache); isCache);
@@ -117,18 +117,6 @@ public class ConfigServiceImpl implements ConfigService {
if (alreadySelectedAsrModelId != null && alreadySelectedAsrModelId.equals(agent.getAsrModelId())) { if (alreadySelectedAsrModelId != null && alreadySelectedAsrModelId.equals(agent.getAsrModelId())) {
agent.setAsrModelId(null); agent.setAsrModelId(null);
} }
String alreadySelectedLlmModelId = (String) selectedModule.get("LLM");
if (alreadySelectedLlmModelId != null && alreadySelectedLlmModelId.equals(agent.getLlmModelId())) {
agent.setLlmModelId(null);
}
String alreadySelectedMemModelId = (String) selectedModule.get("Memory");
if (alreadySelectedMemModelId != null && alreadySelectedMemModelId.equals(agent.getMemModelId())) {
agent.setMemModelId(null);
}
String alreadySelectedIntentModelId = (String) selectedModule.get("Intent");
if (alreadySelectedIntentModelId != null && alreadySelectedIntentModelId.equals(agent.getIntentModelId())) {
agent.setIntentModelId(null);
}
// 构建模块配置 // 构建模块配置
buildModuleConfig( buildModuleConfig(
@@ -269,7 +257,8 @@ public class ConfigServiceImpl implements ConfigService {
if (intentLLMModelId != null && intentLLMModelId.equals(llmModelId)) { if (intentLLMModelId != null && intentLLMModelId.equals(llmModelId)) {
intentLLMModelId = null; intentLLMModelId = null;
} }
} else if ("function_call".equals(map.get("type"))) { }
if (map.get("functions") != null) {
String functionStr = (String) map.get("functions"); String functionStr = (String) map.get("functions");
if (StringUtils.isNotBlank(functionStr)) { if (StringUtils.isNotBlank(functionStr)) {
String[] functions = functionStr.split("\\;"); String[] functions = functionStr.split("\\;");
@@ -1,5 +1,7 @@
package xiaozhi.modules.device.dao; package xiaozhi.modules.device.dao;
import java.util.Date;
import org.apache.ibatis.annotations.Mapper; import org.apache.ibatis.annotations.Mapper;
import com.baomidou.mybatisplus.core.mapper.BaseMapper; import com.baomidou.mybatisplus.core.mapper.BaseMapper;
@@ -8,4 +10,12 @@ import xiaozhi.modules.device.entity.DeviceEntity;
@Mapper @Mapper
public interface DeviceDao extends BaseMapper<DeviceEntity> { public interface DeviceDao extends BaseMapper<DeviceEntity> {
/**
* 获取此智能体全部设备的最后连接时间
*
* @param agentId 智能体id
* @return
*/
Date getAllLastConnectedAtByAgentId(String agentId);
} }
@@ -1,5 +1,6 @@
package xiaozhi.modules.device.service; package xiaozhi.modules.device.service;
import java.util.Date;
import java.util.List; import java.util.List;
import xiaozhi.common.page.PageData; import xiaozhi.common.page.PageData;
@@ -78,4 +79,13 @@ public interface DeviceService extends BaseService<DeviceEntity> {
* @return 激活码 * @return 激活码
*/ */
String geCodeByDeviceId(String deviceId); String geCodeByDeviceId(String deviceId);
/**
* 获取这个智能体设备理的最近的最后连接时间
* @param agentId 智能体id
* @return 返回设备最近的最后连接时间
*/
Date getLatestLastConnectionTime(String agentId);
} }
@@ -118,7 +118,8 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
DeviceEntity deviceById = getDeviceByMacAddress(macAddress); DeviceEntity deviceById = getDeviceByMacAddress(macAddress);
if (deviceById == null || deviceById.getAutoUpdate() != 0) { // 只有在设备已绑定且autoUpdate不为0的情况下才返回固件升级信息
if (deviceById != null && deviceById.getAutoUpdate() != 0) {
String type = deviceReport.getBoard() == null ? null : deviceReport.getBoard().getType(); String type = deviceReport.getBoard() == null ? null : deviceReport.getBoard().getType();
DeviceReportRespDTO.Firmware firmware = buildFirmwareInfo(type, DeviceReportRespDTO.Firmware firmware = buildFirmwareInfo(type,
deviceReport.getApplication() == null ? null : deviceReport.getApplication().getVersion()); deviceReport.getApplication() == null ? null : deviceReport.getApplication().getVersion());
@@ -258,6 +259,20 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
return null; return null;
} }
@Override
public Date getLatestLastConnectionTime(String agentId) {
// 查询是否有缓存时间,有则返回
Date cachedDate = (Date) redisUtils.get(RedisKeys.getAgentDeviceLastConnectedAtById(agentId));
if (cachedDate != null) {
return cachedDate;
}
Date maxDate = deviceDao.getAllLastConnectedAtByAgentId(agentId);
if (maxDate != null) {
redisUtils.set(RedisKeys.getAgentDeviceLastConnectedAtById(agentId), maxDate);
}
return maxDate;
}
private String getDeviceCacheKey(String deviceId) { private String getDeviceCacheKey(String deviceId) {
String safeDeviceId = deviceId.replace(":", "_").toLowerCase(); String safeDeviceId = deviceId.replace(":", "_").toLowerCase();
String dataKey = String.format("ota:activation:data:%s", safeDeviceId); String dataKey = String.format("ota:activation:data:%s", safeDeviceId);
@@ -1,6 +1,9 @@
package xiaozhi.modules.security.config; package xiaozhi.modules.security.config;
import jakarta.servlet.Filter; import java.util.HashMap;
import java.util.LinkedHashMap;
import java.util.Map;
import org.apache.shiro.mgt.SecurityManager; import org.apache.shiro.mgt.SecurityManager;
import org.apache.shiro.session.mgt.SessionManager; import org.apache.shiro.session.mgt.SessionManager;
import org.apache.shiro.spring.LifecycleBeanPostProcessor; import org.apache.shiro.spring.LifecycleBeanPostProcessor;
@@ -11,15 +14,13 @@ import org.apache.shiro.web.mgt.DefaultWebSecurityManager;
import org.apache.shiro.web.session.mgt.DefaultWebSessionManager; import org.apache.shiro.web.session.mgt.DefaultWebSessionManager;
import org.springframework.context.annotation.Bean; import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration; import org.springframework.context.annotation.Configuration;
import jakarta.servlet.Filter;
import xiaozhi.modules.security.oauth2.Oauth2Filter; import xiaozhi.modules.security.oauth2.Oauth2Filter;
import xiaozhi.modules.security.oauth2.Oauth2Realm; import xiaozhi.modules.security.oauth2.Oauth2Realm;
import xiaozhi.modules.security.secret.ServerSecretFilter; import xiaozhi.modules.security.secret.ServerSecretFilter;
import xiaozhi.modules.sys.service.SysParamsService; import xiaozhi.modules.sys.service.SysParamsService;
import java.util.HashMap;
import java.util.LinkedHashMap;
import java.util.Map;
/** /**
* Shiro的配置文件 * Shiro的配置文件
* Copyright (c) 人人开源 All rights reserved. * Copyright (c) 人人开源 All rights reserved.
@@ -85,6 +86,7 @@ public class ShiroConfig {
// 将config路径使用server服务过滤器 // 将config路径使用server服务过滤器
filterMap.put("/config/**", "server"); filterMap.put("/config/**", "server");
filterMap.put("/agent/chat-history/report", "server"); filterMap.put("/agent/chat-history/report", "server");
filterMap.put("/agent/play/**", "anon");
filterMap.put("/**", "oauth2"); filterMap.put("/**", "oauth2");
shiroFilter.setFilterChainDefinitionMap(filterMap); shiroFilter.setFilterChainDefinitionMap(filterMap);
@@ -1,5 +1,6 @@
package xiaozhi.modules.security.config; package xiaozhi.modules.security.config;
import java.text.SimpleDateFormat;
import java.util.List; import java.util.List;
import java.util.TimeZone; import java.util.TimeZone;
@@ -18,6 +19,15 @@ import com.fasterxml.jackson.databind.DeserializationFeature;
import com.fasterxml.jackson.databind.ObjectMapper; import com.fasterxml.jackson.databind.ObjectMapper;
import com.fasterxml.jackson.databind.module.SimpleModule; import com.fasterxml.jackson.databind.module.SimpleModule;
import com.fasterxml.jackson.databind.ser.std.ToStringSerializer; import com.fasterxml.jackson.databind.ser.std.ToStringSerializer;
import com.fasterxml.jackson.datatype.jsr310.JavaTimeModule;
import com.fasterxml.jackson.datatype.jsr310.deser.LocalDateDeserializer;
import com.fasterxml.jackson.datatype.jsr310.deser.LocalDateTimeDeserializer;
import com.fasterxml.jackson.datatype.jsr310.deser.LocalTimeDeserializer;
import com.fasterxml.jackson.datatype.jsr310.ser.LocalDateSerializer;
import com.fasterxml.jackson.datatype.jsr310.ser.LocalDateTimeSerializer;
import com.fasterxml.jackson.datatype.jsr310.ser.LocalTimeSerializer;
import xiaozhi.common.utils.DateUtils;
@Configuration @Configuration
public class WebMvcConfig implements WebMvcConfigurer { public class WebMvcConfig implements WebMvcConfigurer {
@@ -53,10 +63,29 @@ public class WebMvcConfig implements WebMvcConfigurer {
// 忽略未知属性 // 忽略未知属性
mapper.configure(DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES, false); mapper.configure(DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES, false);
// 日期格式转换 // 设置时区
// mapper.setDateFormat(new SimpleDateFormat(DateUtils.DATE_TIME_PATTERN));
mapper.setTimeZone(TimeZone.getTimeZone("GMT+8")); mapper.setTimeZone(TimeZone.getTimeZone("GMT+8"));
// 配置Java8日期时间序列化
JavaTimeModule javaTimeModule = new JavaTimeModule();
javaTimeModule.addSerializer(java.time.LocalDateTime.class, new LocalDateTimeSerializer(
java.time.format.DateTimeFormatter.ofPattern(DateUtils.DATE_TIME_PATTERN)));
javaTimeModule.addSerializer(java.time.LocalDate.class, new LocalDateSerializer(
java.time.format.DateTimeFormatter.ofPattern(DateUtils.DATE_PATTERN)));
javaTimeModule.addSerializer(java.time.LocalTime.class,
new LocalTimeSerializer(java.time.format.DateTimeFormatter.ofPattern("HH:mm:ss")));
javaTimeModule.addDeserializer(java.time.LocalDateTime.class, new LocalDateTimeDeserializer(
java.time.format.DateTimeFormatter.ofPattern(DateUtils.DATE_TIME_PATTERN)));
javaTimeModule.addDeserializer(java.time.LocalDate.class, new LocalDateDeserializer(
java.time.format.DateTimeFormatter.ofPattern(DateUtils.DATE_PATTERN)));
javaTimeModule.addDeserializer(java.time.LocalTime.class,
new LocalTimeDeserializer(java.time.format.DateTimeFormatter.ofPattern("HH:mm:ss")));
mapper.registerModule(javaTimeModule);
// 配置java.util.Date的序列化和反序列化
SimpleDateFormat dateFormat = new SimpleDateFormat(DateUtils.DATE_TIME_PATTERN);
mapper.setDateFormat(dateFormat);
// Long类型转String类型 // Long类型转String类型
SimpleModule simpleModule = new SimpleModule(); SimpleModule simpleModule = new SimpleModule();
simpleModule.addSerializer(Long.class, ToStringSerializer.instance); simpleModule.addSerializer(Long.class, ToStringSerializer.instance);
@@ -96,7 +96,7 @@ public class SysDictDataController {
@GetMapping("/type/{dictType}") @GetMapping("/type/{dictType}")
@Operation(summary = "获取字典数据列表") @Operation(summary = "获取字典数据列表")
@RequiresPermissions("sys:role:superAdmin") @RequiresPermissions("sys:role:normal")
public Result<List<SysDictDataItem>> getDictDataByType(@PathVariable("dictType") String dictType) { public Result<List<SysDictDataItem>> getDictDataByType(@PathVariable("dictType") String dictType) {
List<SysDictDataItem> list = sysDictDataService.getDictDataByType(dictType); List<SysDictDataItem> list = sysDictDataService.getDictDataByType(dictType);
return new Result<List<SysDictDataItem>>().ok(list); return new Result<List<SysDictDataItem>>().ok(list);
@@ -39,7 +39,7 @@ public class SysParamsDTO implements Serializable {
@Schema(description = "值类型") @Schema(description = "值类型")
@NotBlank(message = "{sysparams.valuetype.require}", groups = DefaultGroup.class) @NotBlank(message = "{sysparams.valuetype.require}", groups = DefaultGroup.class)
@Pattern(regexp = "^(string|number|boolean|array)$", message = "{sysparams.valuetype.pattern}", groups = DefaultGroup.class) @Pattern(regexp = "^(string|number|boolean|array|json)$", message = "{sysparams.valuetype.pattern}", groups = DefaultGroup.class)
private String valueType; private String valueType;
@Schema(description = "备注") @Schema(description = "备注")
@@ -127,6 +127,12 @@ public class SysParamsServiceImpl extends BaseServiceImpl<SysParamsDao, SysParam
break; break;
case "json": case "json":
try { try {
// 首先检查是否以 { 开头,以 } 结尾
String trimmedValue = paramValue.trim();
if (!trimmedValue.startsWith("{") || !trimmedValue.endsWith("}")) {
throw new RenException(ErrorCode.PARAM_JSON_INVALID);
}
// 然后尝试解析JSON
JsonUtils.parseObject(paramValue, Object.class); JsonUtils.parseObject(paramValue, Object.class);
} catch (Exception e) { } catch (Exception e) {
throw new RenException(ErrorCode.PARAM_JSON_INVALID); throw new RenException(ErrorCode.PARAM_JSON_INVALID);
@@ -8,6 +8,7 @@ import java.util.regex.Pattern;
import org.apache.commons.lang3.StringUtils; import org.apache.commons.lang3.StringUtils;
import org.slf4j.Logger; import org.slf4j.Logger;
import org.slf4j.LoggerFactory; import org.slf4j.LoggerFactory;
import org.springframework.web.socket.WebSocketHttpHeaders;
import org.springframework.web.socket.client.WebSocketClient; import org.springframework.web.socket.client.WebSocketClient;
import org.springframework.web.socket.client.standard.StandardWebSocketClient; import org.springframework.web.socket.client.standard.StandardWebSocketClient;
@@ -45,8 +46,9 @@ public class WebSocketValidator {
try { try {
WebSocketClient client = new StandardWebSocketClient(); WebSocketClient client = new StandardWebSocketClient();
CompletableFuture<Boolean> future = new CompletableFuture<>(); CompletableFuture<Boolean> future = new CompletableFuture<>();
WebSocketHttpHeaders headers = new WebSocketHttpHeaders();
client.doHandshake(new WebSocketTestHandler(future), null, URI.create(url)); client.execute(new WebSocketTestHandler(future), headers, URI.create(url));
// 等待最多5秒获取连接结果 // 等待最多5秒获取连接结果
return future.get(5, TimeUnit.SECONDS); return future.get(5, TimeUnit.SECONDS);
@@ -1,5 +1,6 @@
-- 初始化智能体聊天记录 -- 初始化智能体聊天记录
DROP TABLE IF EXISTS ai_chat_history; DROP TABLE IF EXISTS ai_chat_history;
DROP TABLE IF EXISTS ai_chat_message;
DROP TABLE IF EXISTS ai_agent_chat_history; DROP TABLE IF EXISTS ai_agent_chat_history;
CREATE TABLE ai_agent_chat_history CREATE TABLE ai_agent_chat_history
( (
@@ -13,7 +14,9 @@ CREATE TABLE ai_agent_chat_history
created_at DATETIME(3) DEFAULT CURRENT_TIMESTAMP(3) NOT NULL COMMENT '创建时间', created_at DATETIME(3) DEFAULT CURRENT_TIMESTAMP(3) NOT NULL COMMENT '创建时间',
updated_at DATETIME(3) DEFAULT CURRENT_TIMESTAMP(3) NOT NULL ON UPDATE CURRENT_TIMESTAMP(3) COMMENT '更新时间', updated_at DATETIME(3) DEFAULT CURRENT_TIMESTAMP(3) NOT NULL ON UPDATE CURRENT_TIMESTAMP(3) COMMENT '更新时间',
INDEX idx_ai_agent_chat_history_mac (mac_address), INDEX idx_ai_agent_chat_history_mac (mac_address),
INDEX idx_ai_agent_chat_history_agent_id (agent_id) INDEX idx_ai_agent_chat_history_session_id (session_id),
INDEX idx_ai_agent_chat_history_agent_id (agent_id),
INDEX idx_ai_agent_chat_history_agent_session_created (agent_id, session_id, created_at)
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT '智能体聊天记录表'; ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT '智能体聊天记录表';
DROP TABLE IF EXISTS ai_agent_chat_audio; DROP TABLE IF EXISTS ai_agent_chat_audio;
@@ -0,0 +1,43 @@
-- 添加百度ASR模型配置
delete from `ai_model_config` where `id` = 'ASR_BaiduASR';
INSERT INTO `ai_model_config` VALUES ('ASR_BaiduASR', 'ASR', 'BaiduASR', '百度语音识别', 0, 1, '{\"type\": \"baidu\", \"app_id\": \"\", \"api_key\": \"\", \"secret_key\": \"\", \"dev_pid\": 1537, \"output_dir\": \"tmp/\"}', NULL, NULL, 7, NULL, NULL, NULL, NULL);
-- 添加百度ASR供应器
delete from `ai_model_provider` where `id` = 'SYSTEM_ASR_BaiduASR';
INSERT INTO `ai_model_provider` (`id`, `model_type`, `provider_code`, `name`, `fields`, `sort`, `creator`, `create_date`, `updater`, `update_date`) VALUES
('SYSTEM_ASR_BaiduASR', 'ASR', 'baidu', '百度语音识别', '[{"key":"app_id","label":"应用AppID","type":"string"},{"key":"api_key","label":"API Key","type":"string"},{"key":"secret_key","label":"Secret Key","type":"string"},{"key":"dev_pid","label":"语言参数","type":"number"},{"key":"output_dir","label":"输出目录","type":"string"}]', 7, 1, NOW(), 1, NOW());
-- 更新百度ASR配置说明
UPDATE `ai_model_config` SET
`doc_link` = 'https://console.bce.baidu.com/ai-engine/old/#/ai/speech/app/list',
`remark` = '百度ASR配置说明:
1. 访问 https://console.bce.baidu.com/ai-engine/old/#/ai/speech/app/list
2. 创建新应用
3. 获取AppID、API Key和Secret Key
4. 填入配置文件中
查看资源额度:https://console.bce.baidu.com/ai-engine/old/#/ai/speech/overview/resource/list
语言参数说明:https://ai.baidu.com/ai-doc/SPEECH/0lbxfnc9b
' WHERE `id` = 'ASR_BaiduASR';
-- 更新豆包供应器字段
update `ai_model_provider` set `fields` =
'[{"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"}]'
where `id` = 'SYSTEM_ASR_DoubaoASR';
-- 更新豆包ASR配置说明
UPDATE `ai_model_config` SET
`doc_link` = 'https://console.volcengine.com/speech/app',
`remark` = '豆包ASR配置说明:
1. 需要在火山引擎控制台创建应用并获取appid和access_token
2. 支持中文语音识别
3. 需要网络连接
4. 输出文件保存在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';
@@ -0,0 +1,6 @@
-- 更新intent_llmM供应器
update `ai_model_provider` set fields = '[{"key":"llm","label":"LLM模型","type":"string"},{"key":"functions","label":"函数列表","type":"dict","dict_name":"functions"}]' where id = 'SYSTEM_Intent_intent_llm';
-- 更新ChatGLMLLM的意图识别配置
update `ai_model_config` set config_json = '{\"type\": \"intent_llm\", \"llm\": \"LLM_ChatGLMLLM\", \"functions\": \"get_weather;get_news_from_newsnow;play_music\"}' where id = 'Intent_intent_llm';
-- 更新函数调用意图识别配置
UPDATE `ai_model_config` SET config_json = REPLACE(config_json, ';get_news;', ';get_news_from_newsnow;') WHERE id = 'Intent_function_call';
@@ -94,9 +94,23 @@ databaseChangeLog:
encoding: utf8 encoding: utf8
path: classpath:db/changelog/202504301341.sql path: classpath:db/changelog/202504301341.sql
- changeSet: - changeSet:
id: 202505012207 id: 202505022134
author: Goody author: Goody
changes: changes:
- sqlFile: - sqlFile:
encoding: utf8 encoding: utf8
path: classpath:db/changelog/202505012207.sql path: classpath:db/changelog/202505022134.sql
- changeSet:
id: 202505081146
author: hrz
changes:
- sqlFile:
encoding: utf8
path: classpath:db/changelog/202505081146.sql
- changeSet:
id: 202505091409
author: hrz
changes:
- sqlFile:
encoding: utf8
path: classpath:db/changelog/202505091409.sql
@@ -21,4 +21,18 @@
id, mac_address, agent_id, session_id, sort, chat_type, content, audio, audio_url, id, mac_address, agent_id, session_id, sort, chat_type, content, audio, audio_url,
created_at, updated_at created_at, updated_at
</sql> </sql>
<delete id="deleteAudioByAgentId">
DELETE FROM ai_agent_chat_audio
WHERE id IN (
SELECT audio_id
FROM ai_agent_chat_history
WHERE agent_id = #{agentId}
);
</delete>
<delete id="deleteHistoryByAgentId">
DELETE FROM ai_agent_chat_history
WHERE agent_id = #{agentId};
</delete>
</mapper> </mapper>
@@ -0,0 +1,12 @@
<?xml version="1.0" encoding="UTF-8"?>
<!DOCTYPE mapper PUBLIC "-//mybatis.org//DTD Mapper 3.0//EN" "http://mybatis.org/dtd/mybatis-3-mapper.dtd">
<mapper namespace="xiaozhi.modules.device.dao.DeviceDao">
<!-- 获取此智能体全部设备的最后连接时间 -->
<select id="getAllLastConnectedAtByAgentId" resultType="java.util.Date">
SELECT last_connected_at FROM ai_device
WHERE
agent_id = #{agentId}
order by
last_connected_at desc limit 0,1
</select>
</mapper>
@@ -1,5 +1,8 @@
package xiaozhi.modules.device; package xiaozhi.modules.device;
import java.util.HashMap;
import java.util.UUID;
import org.junit.jupiter.api.Assertions; import org.junit.jupiter.api.Assertions;
import org.junit.jupiter.api.DisplayName; import org.junit.jupiter.api.DisplayName;
import org.junit.jupiter.api.Test; import org.junit.jupiter.api.Test;
@@ -8,8 +11,9 @@ import org.springframework.boot.test.context.SpringBootTest;
import org.springframework.test.context.ActiveProfiles; import org.springframework.test.context.ActiveProfiles;
import lombok.extern.slf4j.Slf4j; import lombok.extern.slf4j.Slf4j;
import xiaozhi.common.redis.RedisKeys;
import xiaozhi.common.redis.RedisUtils; import xiaozhi.common.redis.RedisUtils;
import xiaozhi.modules.sys.dto.SysUserDTO;
import xiaozhi.modules.sys.service.SysUserService;
@Slf4j @Slf4j
@SpringBootTest @SpringBootTest
@@ -19,18 +23,37 @@ public class DeviceTest {
@Autowired @Autowired
private RedisUtils redisUtils; private RedisUtils redisUtils;
@Autowired
private SysUserService sysUserService;
@Test
public void testSaveUser() {
SysUserDTO userDTO = new SysUserDTO();
userDTO.setUsername("test");
userDTO.setPassword(UUID.randomUUID().toString());
sysUserService.save(userDTO);
}
@Test @Test
@DisplayName("测试写入设备信息") @DisplayName("测试写入设备信息")
public void testWriteDeviceInfo() { public void testWriteDeviceInfo() {
log.info("开始测试写入设备信息..."); log.info("开始测试写入设备信息...");
// 模拟设备MAC地址 // 模拟设备MAC地址
String macAddress = "00:11:22:33:44:55"; String macAddress = "00:11:22:33:44:66";
// 模拟设备验证码 // 模拟设备验证码
String deviceCode = "123456"; String deviceCode = "123456";
String redisKey = RedisKeys.getDeviceCaptchaKey(deviceCode); HashMap<String, Object> map = new HashMap<>();
map.put("mac_address", macAddress);
map.put("activation_code", deviceCode);
map.put("board", "硬件型号");
map.put("app_version", "0.3.13");
String safeDeviceId = macAddress.replace(":", "_").toLowerCase();
String cacheDeviceKey = String.format("ota:activation:data:%s", safeDeviceId);
redisUtils.set(cacheDeviceKey, map, 300);
String redisKey = "ota:activation:code:" + deviceCode;
log.info("Redis Key: {}", redisKey); log.info("Redis Key: {}", redisKey);
// 将设备信息写入Redis // 将设备信息写入Redis
+46
View File
@@ -97,4 +97,50 @@ export default {
}); });
}).send(); }).send();
}, },
// 获取智能体会话列表
getAgentSessions(agentId, params, callback) {
RequestService.sendRequest()
.url(`${getServiceUrl()}/agent/${agentId}/sessions`)
.method('GET')
.data(params)
.success((res) => {
RequestService.clearRequestTime();
callback(res);
})
.fail(() => {
RequestService.reAjaxFun(() => {
this.getAgentSessions(agentId, params, callback);
});
}).send();
},
// 获取智能体聊天记录
getAgentChatHistory(agentId, sessionId, callback) {
RequestService.sendRequest()
.url(`${getServiceUrl()}/agent/${agentId}/chat-history/${sessionId}`)
.method('GET')
.success((res) => {
RequestService.clearRequestTime();
callback(res);
})
.fail(() => {
RequestService.reAjaxFun(() => {
this.getAgentChatHistory(agentId, sessionId, callback);
});
}).send();
},
// 获取音频下载ID
getAudioId(audioId, callback) {
RequestService.sendRequest()
.url(`${getServiceUrl()}/agent/audio/${audioId}`)
.method('POST')
.success((res) => {
RequestService.clearRequestTime();
callback(res);
})
.fail(() => {
RequestService.reAjaxFun(() => {
this.getAudioId(audioId, callback);
});
}).send();
},
} }
Binary file not shown.

After

Width:  |  Height:  |  Size: 11 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 10 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 11 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 10 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 8.1 KiB

@@ -1,5 +1,5 @@
<template> <template>
<el-dialog :visible="visible" @close="handleClose" width="400px" center> <el-dialog :visible="visible" @close="handleClose" width="24%" center>
<div <div
style="margin: 0 10px 10px;display: flex;align-items: center;gap: 10px;font-weight: 700;font-size: 20px;text-align: left;color: #3d4566;"> style="margin: 0 10px 10px;display: flex;align-items: center;gap: 10px;font-weight: 700;font-size: 20px;text-align: left;color: #3d4566;">
<div <div
@@ -1,5 +1,5 @@
<template> <template>
<el-dialog :visible="dialogVisible" @update:visible="handleVisibleChange" width="975px" center <el-dialog :visible="dialogVisible" @update:visible="handleVisibleChange" width="57%" center
custom-class="custom-dialog" :show-close="false" class="center-dialog"> custom-class="custom-dialog" :show-close="false" class="center-dialog">
<div style="margin: 0 18px; text-align: left; padding: 10px; border-radius: 10px;"> <div style="margin: 0 18px; text-align: left; padding: 10px; border-radius: 10px;">
<div style="font-size: 30px; color: #3d4566; margin-top: -10px; margin-bottom: 10px; text-align: center;"> <div style="font-size: 30px; color: #3d4566; margin-top: -10px; margin-bottom: 10px; text-align: center;">
@@ -54,7 +54,7 @@
</el-form-item> </el-form-item>
<el-form-item label="备注" prop="remark" class="prop-remark"> <el-form-item label="备注" prop="remark" class="prop-remark">
<el-input v-model="formData.remark" type="textarea" :rows="3" placeholder="请输入模型备注" <el-input v-model="formData.remark" type="textarea" :rows="3" placeholder="请输入模型备注" :autosize="{ minRows: 3, maxRows: 5 }"
class="custom-input-bg"></el-input> class="custom-input-bg"></el-input>
</el-form-item> </el-form-item>
</el-form> </el-form>
@@ -271,7 +271,7 @@ export default {
} }
.center-dialog .el-dialog { .center-dialog .el-dialog {
margin: 4% 0 auto !important; margin: 0 0 auto !important;
display: flex; display: flex;
flex-direction: column; flex-direction: column;
} }
@@ -1,5 +1,5 @@
<template> <template>
<el-dialog :visible="visible" @close="handleClose" width="400px" center @open="handleOpen"> <el-dialog :visible="visible" @close="handleClose" width="25%" center @open="handleOpen">
<div <div
style="margin: 0 10px 10px;display: flex;align-items: center;gap: 10px;font-weight: 700;font-size: 20px;text-align: left;color: #3d4566;"> style="margin: 0 10px 10px;display: flex;align-items: center;gap: 10px;font-weight: 700;font-size: 20px;text-align: left;color: #3d4566;">
<div <div
@@ -10,7 +10,7 @@
</div> </div>
<div style="height: 1px;background: #e8f0ff;" /> <div style="height: 1px;background: #e8f0ff;" />
<div style="margin: 22px 15px;"> <div style="margin: 22px 15px;">
<div style="font-weight: 400;font-size: 14px;text-align: left;color: #3d4566;"> <div style="font-weight: 400;text-align: left;color: #3d4566;">
<div style="color: red;display: inline-block;">*</div> 智能体名称 <div style="color: red;display: inline-block;">*</div> 智能体名称
</div> </div>
<div class="input-46" style="margin-top: 12px;"> <div class="input-46" style="margin-top: 12px;">
@@ -1,6 +1,6 @@
<template> <template>
<form> <form>
<el-dialog :visible.sync="value" width="400px" center> <el-dialog :visible.sync="dialogVisible" width="24%" center>
<div <div
style="margin: 0 10px 10px;display: flex;align-items: center;gap: 10px;font-weight: 700;font-size: 20px;text-align: left;color: #3d4566;"> style="margin: 0 10px 10px;display: flex;align-items: center;gap: 10px;font-weight: 700;font-size: 20px;text-align: left;color: #3d4566;">
<div <div
@@ -60,11 +60,20 @@ export default {
}, },
data() { data() {
return { return {
dialogVisible: this.value,
oldPassword: "", oldPassword: "",
newPassword: "", newPassword: "",
confirmNewPassword: "" confirmNewPassword: ""
} }
}, },
watch: {
value(val) {
this.dialogVisible = val;
},
dialogVisible(val) {
this.$emit('input', val);
}
},
methods: { methods: {
...mapActions(['logout']), // 引入Vuex的logout action ...mapActions(['logout']), // 引入Vuex的logout action
confirm() { confirm() {
@@ -101,7 +110,7 @@ export default {
this.$emit('input', false); this.$emit('input', false);
}, },
cancel() { cancel() {
this.$emit('input', false); this.dialogVisible = false;
this.resetForm(); this.resetForm();
}, },
resetForm() { resetForm() {
@@ -0,0 +1,447 @@
<template>
<el-dialog :title="'与' + agentName + '的聊天记录' + (currentMacAddress ? '[' + currentMacAddress + ']' : '')"
:visible.sync="dialogVisible" width="80%" :before-close="handleClose" custom-class="chat-history-dialog">
<div class="chat-container">
<div class="session-list" @scroll="handleScroll">
<div v-for="session in sessions" :key="session.sessionId" class="session-item"
:class="{ active: currentSessionId === session.sessionId }" @click="selectSession(session)">
<img :src="getUserAvatar(session.sessionId)" class="avatar" />
<div class="session-info">
<div class="session-time">{{ formatTime(session.createdAt) }}</div>
<div class="message-count">{{ session.chatCount > 99 ? '99' : session.chatCount }}</div>
</div>
</div>
<div v-if="loading" class="loading">加载中...</div>
<div v-if="!hasMore" class="no-more">没有更多记录了</div>
</div>
<div class="chat-content">
<div v-if="currentSessionId" class="messages">
<div v-for="(message, index) in messagesWithTime" :key="message.id">
<div v-if="message.type === 'time'" class="time-divider">
{{ message.content }}
</div>
<div v-else class="message-item" :class="{ 'user-message': message.chatType === 1 }">
<img :src="message.chatType === 1 ? getUserAvatar(currentSessionId) : require('@/assets/xiaozhi-logo.png')"
class="avatar" />
<div class="message-content">
{{ message.content }}
<i v-if="message.audioId" :class="getAudioIconClass(message)"
@click="playAudio(message)" class="audio-icon"></i>
</div>
</div>
</div>
</div>
<div v-else class="no-session-selected">
请选择会话查看聊天记录
</div>
</div>
</div>
</el-dialog>
</template>
<script>
import Api from '@/apis/api';
export default {
name: 'ChatHistoryDialog',
props: {
visible: {
type: Boolean,
default: false
},
agentId: {
type: String,
required: true
},
agentName: {
type: String,
required: true
}
},
data() {
return {
dialogVisible: false,
sessions: [],
messages: [],
currentSessionId: '',
currentMacAddress: '',
page: 1,
limit: 20,
loading: false,
hasMore: true,
scrollTimer: null,
isFirstLoad: true,
playingAudioId: null,
audioElement: null
};
},
watch: {
visible(val) {
this.dialogVisible = val;
if (val) {
this.resetData();
this.loadSessions();
}
},
dialogVisible(val) {
if (!val) {
this.$emit('update:visible', false);
}
}
},
computed: {
messagesWithTime() {
if (!this.messages || this.messages.length === 0) return [];
const result = [];
const TIME_INTERVAL = 60 * 1000; // 1分钟的时间间隔(毫秒)
// 添加第一条消息的时间标记
if (this.messages[0]) {
result.push({
type: 'time',
content: this.formatTime(this.messages[0].createdAt),
id: `time-${Date.now()}-${Math.random().toString(36).substr(2, 9)}`
});
}
// 处理消息列表
for (let i = 0; i < this.messages.length; i++) {
const currentMessage = this.messages[i];
result.push(currentMessage);
// 检查是否需要添加时间标记
if (i < this.messages.length - 1) {
const currentTime = new Date(currentMessage.createdAt).getTime();
const nextTime = new Date(this.messages[i + 1].createdAt).getTime();
if (nextTime - currentTime > TIME_INTERVAL) {
result.push({
type: 'time',
content: this.formatTime(this.messages[i + 1].createdAt),
id: `time-${Date.now()}-${Math.random().toString(36).substr(2, 9)}`
});
}
}
}
return result;
}
},
methods: {
resetData() {
this.sessions = [];
this.messages = [];
this.currentSessionId = '';
this.currentMacAddress = '';
this.page = 1;
this.loading = false;
this.hasMore = true;
this.isFirstLoad = true;
},
handleClose() {
this.dialogVisible = false;
},
loadSessions() {
if (this.loading || (!this.isFirstLoad && !this.hasMore)) {
return;
}
this.loading = true;
const params = {
page: this.page,
limit: this.limit
};
Api.agent.getAgentSessions(this.agentId, params, (res) => {
if (res.data && res.data.data && Array.isArray(res.data.data.list)) {
const list = res.data.data.list;
this.hasMore = list.length === this.limit;
this.sessions = [...this.sessions, ...list];
this.page++;
if (this.sessions.length > 0 && !this.currentSessionId) {
this.selectSession(this.sessions[0]);
}
}
this.loading = false;
this.isFirstLoad = false;
});
},
selectSession(session) {
this.currentSessionId = session.sessionId;
Api.agent.getAgentChatHistory(this.agentId, session.sessionId, (res) => {
if (res.data && res.data.data) {
this.messages = res.data.data;
if (this.messages.length > 0 && this.messages[0].macAddress) {
this.currentMacAddress = this.messages[0].macAddress;
}
}
});
},
handleScroll(e) {
if (this.scrollTimer) {
clearTimeout(this.scrollTimer);
}
this.scrollTimer = setTimeout(() => {
const { scrollTop, scrollHeight, clientHeight } = e.target;
// 当滚动到底部时加载更多
if (scrollHeight - scrollTop <= clientHeight + 50) {
this.loadSessions();
}
}, 200);
},
formatTime(timestamp) {
const date = new Date(timestamp);
const now = new Date();
const today = new Date(now.getFullYear(), now.getMonth(), now.getDate());
const yesterday = new Date(today);
yesterday.setDate(yesterday.getDate() - 1);
const hours = date.getHours().toString().padStart(2, '0');
const minutes = date.getMinutes().toString().padStart(2, '0');
if (date >= today) {
return `今天 ${hours}:${minutes}`;
} else if (date >= yesterday) {
return `昨天 ${hours}:${minutes}`;
} else {
const year = date.getFullYear();
const month = (date.getMonth() + 1).toString().padStart(2, '0');
const day = date.getDate().toString().padStart(2, '0');
return `${year}-${month}-${day} ${hours}:${minutes}`;
}
},
getAudioIconClass(message) {
if (this.playingAudioId === message.audioId) {
return 'el-icon-loading';
}
return 'el-icon-video-play';
},
playAudio(message) {
if (this.playingAudioId === message.audioId) {
// 如果正在播放当前音频,则停止播放
if (this.audioElement) {
this.audioElement.pause();
this.audioElement = null;
}
this.playingAudioId = null;
return;
}
// 停止当前正在播放的音频
if (this.audioElement) {
this.audioElement.pause();
this.audioElement = null;
}
// 先获取音频下载ID
this.playingAudioId = message.audioId;
Api.agent.getAudioId(message.audioId, (res) => {
if (res.data && res.data.data) {
// 使用获取到的下载ID播放音频
this.audioElement = new Audio(Api.getServiceUrl() + `/agent/play/${res.data.data}`);
this.audioElement.onended = () => {
this.playingAudioId = null;
this.audioElement = null;
};
this.audioElement.play();
}
});
},
getUserAvatar(sessionId) {
// 从 sessionId 中提取所有数字
const numbers = sessionId.match(/\d+/g);
if (!numbers) return require('@/assets/user-avatar1.png');
// 将所有数字相加
const sum = numbers.reduce((acc, num) => acc + parseInt(num), 0);
// 计算模5并加1,得到1-5之间的数字
const avatarIndex = (sum % 5) + 1;
// 返回对应的头像图片
return require(`@/assets/user-avatar${avatarIndex}.png`);
}
}
};
</script>
<style scoped>
.chat-container {
display: flex;
height: 100%;
}
.session-list {
width: 250px;
border-right: 1px solid #eee;
overflow-y: auto;
padding: 10px;
}
.session-item {
display: flex;
align-items: center;
padding: 10px;
cursor: pointer;
border-radius: 8px;
margin-bottom: 10px;
}
.session-item:hover {
background-color: #f5f5f5;
}
.session-item.active {
background-color: #e6f7ff;
}
.avatar {
width: 40px;
height: 40px;
border-radius: 50%;
margin-right: 10px;
}
.session-info {
flex: 1;
}
.session-time {
font-size: 14px;
color: #272727;
float: left;
height: 30px;
line-height: 30px;
width: calc(100% - 30px);
/* 为消息数量留出空间 */
overflow: hidden;
text-overflow: ellipsis;
white-space: nowrap;
}
.message-count {
font-size: 14px;
color: #fff;
background-color: #b4b4b4;
border-radius: 20px;
float: left;
width: 20px;
height: 20px;
line-height: 20px;
margin-top: 5px;
margin-left: 5px;
}
.chat-content {
flex: 1;
padding: 20px;
overflow-y: auto;
}
.message-item {
display: flex;
margin-bottom: 20px;
}
.message-item.user-message {
flex-direction: row-reverse;
}
.message-content {
max-width: 60%;
padding: 10px 15px;
border-radius: 8px;
background-color: #f0f0f0;
margin: 0 10px;
text-align: left;
line-height: 20px;
position: relative;
display: flex;
align-items: center;
}
.audio-icon {
font-size: 20px;
cursor: pointer;
margin: 0 5px;
color: #1890ff;
}
.user-message .message-content {
background-color: #1890ff;
color: white;
flex-direction: row-reverse;
}
.user-message .audio-icon {
color: white;
}
.loading,
.no-more {
text-align: center;
padding: 10px 10px 30px 10px;
color: #999;
}
.no-session-selected {
display: flex;
justify-content: center;
align-items: center;
height: 100%;
color: #999;
}
.time-divider {
text-align: center;
margin: 10px 0;
color: #999;
font-size: 12px;
}
.time-divider::before,
.time-divider::after {
content: '';
display: inline-block;
width: 30%;
height: 1px;
background-color: #eee;
vertical-align: middle;
margin: 0 10px;
}
</style>
<style>
.chat-history-dialog {
display: flex;
flex-direction: column;
min-width: 700px;
margin: 0 !important;
position: absolute;
top: 50%;
left: 50%;
transform: translate(-50%, -50%);
height: 90vh;
max-width: 85vw;
border-radius: 12px;
overflow: hidden;
}
.chat-history-dialog .el-dialog__header {
background-color: #e6f7ff;
padding: 15px 20px;
}
.chat-history-dialog .el-dialog__body {
padding: 0;
overflow: hidden;
height: calc(90vh - 54px);
/* 减去标题栏的高度 */
}
</style>
+28 -1
View File
@@ -26,9 +26,12 @@
<div class="settings-btn" @click="handleDeviceManage"> <div class="settings-btn" @click="handleDeviceManage">
设备管理({{ device.deviceCount }}) 设备管理({{ device.deviceCount }})
</div> </div>
<div class="settings-btn" @click="handleChatHistory">
聊天记录
</div>
</div> </div>
<div class="version-info"> <div class="version-info">
<div>最近对话{{ device.lastConnectedAt }}</div> <div>最近对话{{ formattedLastConnectedTime }}</div>
</div> </div>
</div> </div>
</template> </template>
@@ -42,6 +45,27 @@ export default {
data() { data() {
return { switchValue: false } return { switchValue: false }
}, },
computed: {
formattedLastConnectedTime() {
if (!this.device.lastConnectedAt) return '暂未对话';
const lastTime = new Date(this.device.lastConnectedAt);
const now = new Date();
const diffMinutes = Math.floor((now - lastTime) / (1000 * 60));
if (diffMinutes <= 1) {
return '刚刚';
} else if (diffMinutes < 60) {
return `${diffMinutes}分钟前`;
} else if (diffMinutes < 24 * 60) {
const hours = Math.floor(diffMinutes / 60);
const minutes = diffMinutes % 60;
return `${hours}小时${minutes > 0 ? minutes + '分钟' : ''}`;
} else {
return this.device.lastConnectedAt;
}
}
},
methods: { methods: {
handleDelete() { handleDelete() {
this.$emit('delete', this.device.agentId) this.$emit('delete', this.device.agentId)
@@ -51,6 +75,9 @@ export default {
}, },
handleDeviceManage() { handleDeviceManage() {
this.$router.push({ path: '/device-management', query: { agentId: this.device.agentId } }); this.$router.push({ path: '/device-management', query: { agentId: this.device.agentId } });
},
handleChatHistory() {
this.$emit('chat-history', { agentId: this.device.agentId, agentName: this.device.agentName })
} }
} }
} }
@@ -1,5 +1,5 @@
<template> <template>
<el-dialog :title="title" :visible.sync="visible" width="500px" @close="handleClose"> <el-dialog :title="title" :visible.sync="dialogVisible" width="30%" @close="handleClose">
<el-form :model="form" :rules="rules" ref="form" label-width="100px"> <el-form :model="form" :rules="rules" ref="form" label-width="100px">
<el-form-item label="字典标签" prop="dictLabel"> <el-form-item label="字典标签" prop="dictLabel">
<el-input v-model="form.dictLabel" placeholder="请输入字典标签"></el-input> <el-input v-model="form.dictLabel" placeholder="请输入字典标签"></el-input>
@@ -41,6 +41,7 @@ export default {
}, },
data() { data() {
return { return {
dialogVisible: this.visible,
form: { form: {
id: null, id: null,
dictTypeId: null, dictTypeId: null,
@@ -70,12 +71,18 @@ export default {
} }
}, },
immediate: true immediate: true
},
visible(val) {
this.dialogVisible = val;
},
dialogVisible(val) {
this.$emit('update:visible', val);
} }
}, },
methods: { methods: {
handleClose() { handleClose() {
this.$emit('update:visible', false) this.dialogVisible = false;
this.resetForm() this.resetForm();
}, },
resetForm() { resetForm() {
this.form = { this.form = {
@@ -102,4 +109,8 @@ export default {
.dialog-footer { .dialog-footer {
text-align: right; text-align: right;
} }
:deep(.el-dialog) {
border-radius: 15px;
}
</style> </style>
@@ -1,5 +1,5 @@
<template> <template>
<el-dialog :title="title" :visible.sync="visible" width="500px" @close="handleClose"> <el-dialog :title="title" :visible.sync="dialogVisible" width="30%" @close="handleClose">
<el-form :model="form" :rules="rules" ref="form" label-width="120px"> <el-form :model="form" :rules="rules" ref="form" label-width="120px">
<el-form-item label="字典类型名称" prop="dictName"> <el-form-item label="字典类型名称" prop="dictName">
<el-input v-model="form.dictName" placeholder="请输入字典类型名称"></el-input> <el-input v-model="form.dictName" placeholder="请输入字典类型名称"></el-input>
@@ -34,6 +34,7 @@ export default {
}, },
data() { data() {
return { return {
dialogVisible: this.visible,
form: { form: {
id: null, id: null,
dictName: '', dictName: '',
@@ -46,6 +47,12 @@ export default {
} }
}, },
watch: { watch: {
visible(val) {
this.dialogVisible = val;
},
dialogVisible(val) {
this.$emit('update:visible', val);
},
dictTypeData: { dictTypeData: {
handler(val) { handler(val) {
if (val) { if (val) {
@@ -57,7 +64,7 @@ export default {
}, },
methods: { methods: {
handleClose() { handleClose() {
this.$emit('update:visible', false) this.dialogVisible = false;
this.resetForm() this.resetForm()
}, },
resetForm() { resetForm() {
@@ -83,4 +90,8 @@ export default {
.dialog-footer { .dialog-footer {
text-align: right; text-align: right;
} }
:deep(.el-dialog) {
border-radius: 15px;
}
</style> </style>
@@ -1,5 +1,5 @@
<template> <template>
<el-dialog :title="title" :visible.sync="visible" width="500px" @close="handleClose" @open="handleOpen"> <el-dialog :title="title" :visible.sync="dialogVisible" width="30%" @close="handleClose" @open="handleOpen">
<el-form ref="form" :model="form" :rules="rules" label-width="100px"> <el-form ref="form" :model="form" :rules="rules" label-width="100px">
<el-form-item label="固件名称" prop="firmwareName"> <el-form-item label="固件名称" prop="firmwareName">
<el-input v-model="form.firmwareName" placeholder="请输入固件名称(板子+版本号)"></el-input> <el-input v-model="form.firmwareName" placeholder="请输入固件名称(板子+版本号)"></el-input>
@@ -59,11 +59,13 @@ export default {
default: () => [] default: () => []
} }
}, },
data() { data() {
return { return {
uploadProgress: 0, uploadProgress: 0,
uploadStatus: '', uploadStatus: '',
isUploading: false, isUploading: false,
dialogVisible: this.visible,
rules: { rules: {
firmwareName: [ firmwareName: [
{ required: true, message: '请输入固件名称(板子+版本号)', trigger: 'blur' } { required: true, message: '请输入固件名称(板子+版本号)', trigger: 'blur' }
@@ -90,10 +92,18 @@ export default {
created() { created() {
// 移除 getDictDataByType 调用 // 移除 getDictDataByType 调用
}, },
watch: {
visible(val) {
this.dialogVisible = val;
},
dialogVisible(val) {
this.$emit('update:visible', val);
},
},
methods: { methods: {
// 移除 getFirmwareTypes 方法 // 移除 getFirmwareTypes 方法
handleClose() { handleClose() {
this.$refs.form.clearValidate(); this.dialogVisible = false;
this.$emit('cancel'); this.$emit('cancel');
}, },
handleCancel() { handleCancel() {
@@ -201,13 +211,17 @@ export default {
</script> </script>
<style lang="scss" scoped> <style lang="scss" scoped>
::v-deep .el-dialog {
border-radius: 20px;
}
.upload-demo { .upload-demo {
text-align: left; text-align: left;
} }
.el-upload__tip { .el-upload__tip {
line-height: 1.2; line-height: 1.2;
padding-top: 5px; padding-top: 2%;
color: #909399; color: #909399;
} }
@@ -1,6 +1,6 @@
<template> <template>
<el-dialog :visible.sync="dialogVisible" width="975px" center custom-class="custom-dialog" :show-close="false" <el-dialog :visible.sync="dialogVisible" width="57%" center custom-class="custom-dialog" :show-close="false"
class="center-dialog"> class="center-dialog" >
<div style="margin: 0 18px; text-align: left; padding: 10px; border-radius: 10px;"> <div style="margin: 0 18px; text-align: left; padding: 10px; border-radius: 10px;">
<div style="font-size: 30px; color: #3d4566; margin-top: -10px; margin-bottom: 10px; text-align: center;"> <div style="font-size: 30px; color: #3d4566; margin-top: -10px; margin-bottom: 10px; text-align: center;">
修改模型 修改模型
@@ -53,7 +53,7 @@
</el-form-item> </el-form-item>
<el-form-item label="备注" prop="remark" class="prop-remark"> <el-form-item label="备注" prop="remark" class="prop-remark">
<el-input v-model="form.remark" type="textarea" :rows="3" placeholder="请输入模型备注" <el-input v-model="form.remark" type="textarea" :rows="3" placeholder="请输入模型备注" :autosize="{ minRows: 3, maxRows: 5 }"
class="custom-input-bg"></el-input> class="custom-input-bg"></el-input>
</el-form-item> </el-form-item>
</el-form> </el-form>
@@ -296,7 +296,7 @@ export default {
}; };
</script> </script>
<style scoped> <style lang="scss" scoped>
.custom-dialog { .custom-dialog {
position: relative; position: relative;
border-radius: 20px; border-radius: 20px;
@@ -316,11 +316,6 @@ export default {
justify-content: center; justify-content: center;
} }
.center-dialog .el-dialog {
margin: 4% 0 auto !important;
display: flex;
flex-direction: column;
}
.custom-close-btn { .custom-close-btn {
position: absolute; position: absolute;
+1 -8
View File
@@ -487,7 +487,7 @@ export default {
}; };
</script> </script>
<style scoped> <style lang="scss" scoped>
::v-deep .el-dialog { ::v-deep .el-dialog {
border-radius: 8px !important; border-radius: 8px !important;
@@ -648,12 +648,6 @@ export default {
margin: 0 auto; margin: 0 auto;
} }
/* 新增按钮组样式 */
.action-buttons {
bottom: 20px;
padding-top: 10px;
}
.action-buttons .el-button { .action-buttons .el-button {
padding: 8px 15px; padding: 8px 15px;
font-size: 11px; font-size: 11px;
@@ -692,7 +686,6 @@ export default {
position: static; position: static;
padding: 15px 0; padding: 15px 0;
background: white; background: white;
box-shadow: 0 -2px 12px rgba(0,0,0,0.05);
} }
/* 输入框自适应 */ /* 输入框自适应 */
@@ -294,20 +294,13 @@ export default {
this.loading = false; this.loading = false;
if (data.code === 0) { if (data.code === 0) {
this.deviceList = data.data.map(device => { this.deviceList = data.data.map(device => {
const bindDate = new Date(device.createDate);
const formattedBindTime = `${bindDate.getFullYear()}-${(bindDate.getMonth() + 1).toString().padStart(2, '0')}-${bindDate.getDate().toString().padStart(2, '0')} ${bindDate.getHours().toString().padStart(2, '0')}:${bindDate.getMinutes().toString().padStart(2, '0')}:${bindDate.getSeconds().toString().padStart(2, '0')}`;
let formattedLastConversation = '';
if (device.lastConnectedAt) {
const lastConvoDate = new Date(device.lastConnectedAt);
formattedLastConversation = `${lastConvoDate.getFullYear()}-${(lastConvoDate.getMonth() + 1).toString().padStart(2, '0')}-${lastConvoDate.getDate().toString().padStart(2, '0')} ${lastConvoDate.getHours().toString().padStart(2, '0')}:${lastConvoDate.getMinutes().toString().padStart(2, '0')}:${lastConvoDate.getSeconds().toString().padStart(2, '0')}`;
}
return { return {
device_id: device.id, device_id: device.id,
model: device.board, model: device.board,
firmwareVersion: device.appVersion, firmwareVersion: device.appVersion,
macAddress: device.macAddress, macAddress: device.macAddress,
bindTime: formattedBindTime, bindTime: device.createDate,
lastConversation: formattedLastConversation, lastConversation: device.lastConnectedAt,
remark: device.alias, remark: device.alias,
isEdit: false, isEdit: false,
otaSwitch: device.autoUpdate === 1, otaSwitch: device.autoUpdate === 1,
+32 -16
View File
@@ -49,9 +49,13 @@
v-loading="dictDataLoading" element-loading-text="拼命加载中" v-loading="dictDataLoading" element-loading-text="拼命加载中"
element-loading-spinner="el-icon-loading" element-loading-spinner="el-icon-loading"
element-loading-background="rgba(255, 255, 255, 0.7)" element-loading-background="rgba(255, 255, 255, 0.7)"
@selection-change="handleDictDataSelectionChange" class="data-table" class="data-table"
header-row-class-name="table-header"> header-row-class-name="table-header">
<el-table-column type="selection" width="55" align="center"></el-table-column> <el-table-column label="选择" align="center" width="55">
<template slot-scope="scope">
<el-checkbox v-model="scope.row.selected"></el-checkbox>
</template>
</el-table-column>
<el-table-column label="字典标签" prop="dictLabel" align="center"></el-table-column> <el-table-column label="字典标签" prop="dictLabel" align="center"></el-table-column>
<el-table-column label="字典值" prop="dictValue" align="center"></el-table-column> <el-table-column label="字典值" prop="dictValue" align="center"></el-table-column>
<el-table-column label="排序" prop="sort" align="center"></el-table-column> <el-table-column label="排序" prop="sort" align="center"></el-table-column>
@@ -153,7 +157,6 @@ export default {
// 字典数据相关 // 字典数据相关
dictDataList: [], dictDataList: [],
dictDataLoading: false, dictDataLoading: false,
selectedDictData: [],
isAllDictDataSelected: false, isAllDictDataSelected: false,
dictDataDialogVisible: false, dictDataDialogVisible: false,
dictDataDialogTitle: '新增字典数据', dictDataDialogTitle: '新增字典数据',
@@ -265,7 +268,10 @@ export default {
dictValue: '' dictValue: ''
}, ({ data }) => { }, ({ data }) => {
if (data.code === 0) { if (data.code === 0) {
this.dictDataList = data.data.list this.dictDataList = data.data.list.map(item => ({
...item,
selected: false
}))
this.total = data.data.total this.total = data.data.total
} else { } else {
this.$message.error(data.msg || '获取字典数据失败') this.$message.error(data.msg || '获取字典数据失败')
@@ -273,16 +279,11 @@ export default {
this.dictDataLoading = false this.dictDataLoading = false
}) })
}, },
handleDictDataSelectionChange(val) {
this.selectedDictData = val
this.isAllDictDataSelected = val.length === this.dictDataList.length
},
selectAllDictData() { selectAllDictData() {
if (this.isAllDictDataSelected) { this.isAllDictDataSelected = !this.isAllDictDataSelected
this.$refs.dictDataTable.clearSelection() this.dictDataList.forEach(row => {
} else { row.selected = this.isAllDictDataSelected
this.$refs.dictDataTable.toggleAllSelection() })
}
}, },
showAddDictDataDialog() { showAddDictDataDialog() {
if (!this.selectedDictType) { if (!this.selectedDictType) {
@@ -329,17 +330,18 @@ export default {
}) })
}, },
batchDeleteDictData() { batchDeleteDictData() {
if (this.selectedDictData.length === 0) { const selectedRows = this.dictDataList.filter(row => row.selected)
if (selectedRows.length === 0) {
this.$message.warning('请选择要删除的字典数据') this.$message.warning('请选择要删除的字典数据')
return return
} }
this.$confirm('确定要删除选中的字典数据吗?', '提示', { this.$confirm(`确定要删除选中的${selectedRows.length}字典数据吗?`, '提示', {
confirmButtonText: '确定', confirmButtonText: '确定',
cancelButtonText: '取消', cancelButtonText: '取消',
type: 'warning' type: 'warning'
}).then(() => { }).then(() => {
const ids = this.selectedDictData.map(item => item.id) const ids = selectedRows.map(item => item.id)
dictApi.deleteDictData(ids, ({ data }) => { dictApi.deleteDictData(ids, ({ data }) => {
if (data.code === 0) { if (data.code === 0) {
this.$message.success('删除成功') this.$message.success('删除成功')
@@ -832,4 +834,18 @@ export default {
flex: 1; flex: 1;
overflow: hidden; overflow: hidden;
} }
:deep(.el-checkbox__inner) {
background-color: #eeeeee !important;
border-color: #cccccc !important;
}
:deep(.el-checkbox__inner:hover) {
border-color: #cccccc !important;
}
:deep(.el-checkbox__input.is-checked .el-checkbox__inner) {
background-color: #5f70f3 !important;
border-color: #5f70f3 !important;
}
</style> </style>
@@ -25,8 +25,19 @@
</template> </template>
</el-table-column> </el-table-column>
<el-table-column label="参数编码" prop="paramCode" align="center"></el-table-column> <el-table-column label="参数编码" prop="paramCode" align="center"></el-table-column>
<el-table-column label="参数值" prop="paramValue" align="center" <el-table-column label="参数值" prop="paramValue" align="center" show-overflow-tooltip>
show-overflow-tooltip></el-table-column> <template slot-scope="scope">
<div v-if="isSensitiveParam(scope.row.paramCode)">
<span v-if="!scope.row.showValue">{{ maskSensitiveValue(scope.row.paramValue)
}}</span>
<span v-else>{{ scope.row.paramValue }}</span>
<el-button size="mini" type="text" @click="toggleSensitiveValue(scope.row)">
{{ scope.row.showValue ? '隐藏' : '查看' }}
</el-button>
</div>
<span v-else>{{ scope.row.paramValue }}</span>
</template>
</el-table-column>
<el-table-column label="备注" prop="remark" align="center"></el-table-column> <el-table-column label="备注" prop="remark" align="center"></el-table-column>
<el-table-column label="操作" align="center"> <el-table-column label="操作" align="center">
<template slot-scope="scope"> <template slot-scope="scope">
@@ -100,6 +111,7 @@ export default {
dialogVisible: false, dialogVisible: false,
dialogTitle: "新增参数", dialogTitle: "新增参数",
isAllSelected: false, isAllSelected: false,
sensitive_keys: ["api_key", "personal_access_token", "access_token", "token", "secret", "access_key_secret", "secret_key"],
paramForm: { paramForm: {
id: null, id: null,
paramCode: "", paramCode: "",
@@ -152,7 +164,8 @@ export default {
if (data.code === 0) { if (data.code === 0) {
this.paramsList = data.data.list.map(item => ({ this.paramsList = data.data.list.map(item => ({
...item, ...item,
selected: false selected: false,
showValue: false
})); }));
this.total = data.data.total; this.total = data.data.total;
} else { } else {
@@ -314,7 +327,18 @@ export default {
goToPage(page) { goToPage(page) {
this.currentPage = page; this.currentPage = page;
this.fetchParams(); this.fetchParams();
} },
isSensitiveParam(paramCode) {
return this.sensitive_keys.some(key => paramCode.toLowerCase().includes(key.toLowerCase()));
},
maskSensitiveValue(value) {
if (!value) return '';
if (value.length <= 8) return '****';
return value.substring(0, 4) + '****' + value.substring(value.length - 4);
},
toggleSensitiveValue(row) {
this.$set(row, 'showValue', !row.showValue);
},
}, },
}; };
</script> </script>
+1 -10
View File
@@ -26,11 +26,7 @@
<el-table-column label="用户Id" prop="userid" align="center"></el-table-column> <el-table-column label="用户Id" prop="userid" align="center"></el-table-column>
<el-table-column label="手机号码" prop="mobile" align="center"></el-table-column> <el-table-column label="手机号码" prop="mobile" align="center"></el-table-column>
<el-table-column label="设备数量" prop="deviceCount" align="center"></el-table-column> <el-table-column label="设备数量" prop="deviceCount" align="center"></el-table-column>
<el-table-column label="注册时间" prop="createDate" align="center"> <el-table-column label="注册时间" prop="createDate" align="center"></el-table-column>
<template slot-scope="scope">
{{ formatDate(scope.row.createDate) }}
</template>
</el-table-column>
<el-table-column label="状态" prop="status" align="center"> <el-table-column label="状态" prop="status" align="center">
<template slot-scope="scope"> <template slot-scope="scope">
<el-tag v-if="scope.row.status === 1" type="success">正常</el-tag> <el-tag v-if="scope.row.status === 1" type="success">正常</el-tag>
@@ -342,11 +338,6 @@ export default {
// 用户取消操作 // 用户取消操作
}); });
}, },
formatDate(dateString) {
if (!dateString) return '';
const date = new Date(dateString);
return `${date.getFullYear()}-${(date.getMonth() + 1).toString().padStart(2, '0')}-${date.getDate().toString().padStart(2, '0')} ${date.getHours().toString().padStart(2, '0')}:${date.getMinutes().toString().padStart(2, '0')}:${date.getSeconds().toString().padStart(2, '0')}`;
},
}, },
}; };
</script> </script>
+21 -21
View File
@@ -32,11 +32,7 @@
</div> </div>
<div class="device-list-container"> <div class="device-list-container">
<template v-if="isLoading"> <template v-if="isLoading">
<div <div v-for="i in skeletonCount" :key="'skeleton-' + i" class="skeleton-item">
v-for="i in skeletonCount"
:key="'skeleton-'+i"
class="skeleton-item"
>
<div class="skeleton-image"></div> <div class="skeleton-image"></div>
<div class="skeleton-content"> <div class="skeleton-content">
<div class="skeleton-line"></div> <div class="skeleton-line"></div>
@@ -46,14 +42,8 @@
</template> </template>
<template v-else> <template v-else>
<DeviceItem <DeviceItem v-for="(item, index) in devices" :key="index" :device="item" @configure="goToRoleConfig"
v-for="(item, index) in devices" @deviceManage="handleDeviceManage" @delete="handleDeleteAgent" @chat-history="handleShowChatHistory" />
:key="index"
:device="item"
@configure="goToRoleConfig"
@deviceManage="handleDeviceManage"
@delete="handleDeleteAgent"
/>
</template> </template>
</div> </div>
</div> </div>
@@ -62,6 +52,7 @@
<el-footer> <el-footer>
<version-footer /> <version-footer />
</el-footer> </el-footer>
<chat-history-dialog :visible.sync="showChatHistory" :agent-id="currentAgentId" :agent-name="currentAgentName" />
</div> </div>
</template> </template>
@@ -69,13 +60,14 @@
<script> <script>
import Api from '@/apis/api'; import Api from '@/apis/api';
import AddWisdomBodyDialog from '@/components/AddWisdomBodyDialog.vue'; import AddWisdomBodyDialog from '@/components/AddWisdomBodyDialog.vue';
import ChatHistoryDialog from '@/components/ChatHistoryDialog.vue';
import DeviceItem from '@/components/DeviceItem.vue'; import DeviceItem from '@/components/DeviceItem.vue';
import HeaderBar from '@/components/HeaderBar.vue'; import HeaderBar from '@/components/HeaderBar.vue';
import VersionFooter from '@/components/VersionFooter.vue'; import VersionFooter from '@/components/VersionFooter.vue';
export default { export default {
name: 'HomePage', name: 'HomePage',
components: { DeviceItem, AddWisdomBodyDialog, HeaderBar, VersionFooter }, components: { DeviceItem, AddWisdomBodyDialog, HeaderBar, VersionFooter, ChatHistoryDialog },
data() { data() {
return { return {
addDeviceDialogVisible: false, addDeviceDialogVisible: false,
@@ -85,6 +77,9 @@ export default {
searchRegex: null, searchRegex: null,
isLoading: true, isLoading: true,
skeletonCount: localStorage.getItem('skeletonCount') || 8, skeletonCount: localStorage.getItem('skeletonCount') || 8,
showChatHistory: false,
currentAgentId: '',
currentAgentName: ''
} }
}, },
@@ -177,6 +172,11 @@ export default {
} }
}); });
}).catch(() => { }); }).catch(() => { });
},
handleShowChatHistory({ agentId, agentName }) {
this.currentAgentId = agentId;
this.currentAgentName = agentName;
this.showChatHistory = true;
} }
} }
} }
@@ -302,7 +302,9 @@ export default {
/* 骨架屏动画 */ /* 骨架屏动画 */
@keyframes shimmer { @keyframes shimmer {
100% { transform: translateX(100%); } 100% {
transform: translateX(100%);
}
} }
.skeleton-item { .skeleton-item {
@@ -353,12 +355,10 @@ export default {
left: 0; left: 0;
width: 50%; width: 50%;
height: 100%; height: 100%;
background: linear-gradient( background: linear-gradient(90deg,
90deg, rgba(255, 255, 255, 0),
rgba(255,255,255,0), rgba(255, 255, 255, 0.3),
rgba(255,255,255,0.3), rgba(255, 255, 255, 0));
rgba(255,255,255,0)
);
animation: shimmer 1.5s infinite; animation: shimmer 1.5s infinite;
} }
</style> </style>
+6 -7
View File
@@ -318,7 +318,6 @@ export default {
<style scoped> <style scoped>
.welcome { .welcome {
min-width: 900px; min-width: 900px;
min-height: 506px;
height: 100vh; height: 100vh;
display: flex; display: flex;
position: relative; position: relative;
@@ -334,7 +333,7 @@ export default {
display: flex; display: flex;
justify-content: space-between; justify-content: space-between;
align-items: center; align-items: center;
padding: 16px 24px; padding: 1.5vh 24px;
} }
.page-title { .page-title {
@@ -344,7 +343,7 @@ export default {
} }
.main-wrapper { .main-wrapper {
margin: 5px 22px; margin: 1vh 22px;
border-radius: 15px; border-radius: 15px;
height: calc(100vh - 24vh); height: calc(100vh - 24vh);
box-shadow: 0 2px 12px rgba(0, 0, 0, 0.1); box-shadow: 0 2px 12px rgba(0, 0, 0, 0.1);
@@ -416,7 +415,7 @@ export default {
} }
.form-content { .form-content {
padding: 20px 0; padding: 2vh 0;
} }
.form-grid { .form-grid {
@@ -450,11 +449,11 @@ export default {
} }
.template-item { .template-item {
height: 37px; height: 4vh;
width: 76px; width: 76px;
border-radius: 8px; border-radius: 8px;
background: #e6ebff; background: #e6ebff;
line-height: 37px; line-height: 4vh;
font-weight: 400; font-weight: 400;
font-size: 11px; font-size: 11px;
text-align: center; text-align: center;
@@ -471,7 +470,7 @@ export default {
display: flex; display: flex;
flex-wrap: wrap; flex-wrap: wrap;
gap: 8px; gap: 8px;
margin-top: 20px; margin-top: 2vh;
align-items: center; align-items: center;
} }
+36 -19
View File
@@ -7,33 +7,47 @@ from core.ota_server import SimpleOtaServer
from core.utils.util import check_ffmpeg_installed from core.utils.util import check_ffmpeg_installed
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 aioconsole import ainput
TAG = __name__ TAG = __name__
logger = setup_logging() logger = setup_logging()
async def wait_for_exit(): async def wait_for_exit() -> None:
"""Windows 和 Linux 兼容的退出监听""" """
阻塞直到收到 CtrlC / SIGTERM。
- Unix: 使用 add_signal_handler
- Windows: 依赖 KeyboardInterrupt
"""
loop = asyncio.get_running_loop() loop = asyncio.get_running_loop()
stop_event = asyncio.Event() stop_event = asyncio.Event()
if sys.platform == "win32": if sys.platform != "win32": # Unix / macOS
# Windows: 用 sys.stdin.read() 监听 Ctrl + C for sig in (signal.SIGINT, signal.SIGTERM):
await loop.run_in_executor(None, sys.stdin.read) loop.add_signal_handler(sig, stop_event.set)
else:
# Linux/macOS: 用 signal 监听 Ctrl + C
def stop():
stop_event.set()
loop.add_signal_handler(signal.SIGINT, stop)
loop.add_signal_handler(signal.SIGTERM, stop) # 支持 kill 进程
await stop_event.wait() await stop_event.wait()
else:
# Windowsawait一个永远pending的fut
# 让 KeyboardInterrupt 冒泡到 asyncio.run,以此消除遗留普通线程导致进程退出阻塞的问题
try:
await asyncio.Future()
except KeyboardInterrupt: # CtrlC
pass
async def monitor_stdin():
"""监控标准输入,消费回车键"""
while True:
await ainput() # 异步等待输入,消费回车
async def main(): async def main():
check_ffmpeg_installed() check_ffmpeg_installed()
config = load_config() config = load_config()
# 添加 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())
@@ -74,19 +88,22 @@ async def main():
) )
try: try:
await wait_for_exit() # 监听退出信号 await wait_for_exit() # 阻塞直到收到退出信号
except asyncio.CancelledError: except asyncio.CancelledError:
print("任务被取消,清理资源中...") print("任务被取消,清理资源中...")
finally: finally:
# 取消所有任务(关键修复点)
stdin_task.cancel()
ws_task.cancel() ws_task.cancel()
if ota_task: if ota_task:
ota_task.cancel() ota_task.cancel()
try:
await ws_task # 等待任务终止(必须加超时)
if ota_task: await asyncio.wait(
await ota_task [stdin_task, ws_task, ota_task] if ota_task else [stdin_task, ws_task],
except asyncio.CancelledError: timeout=3.0,
pass return_when=asyncio.ALL_COMPLETED
)
print("服务器已关闭,程序退出。") print("服务器已关闭,程序退出。")
+25 -4
View File
@@ -104,12 +104,13 @@ plugins:
get_weather: { "api_key": "a861d0d5e7bf4ee1a83d9a9e4f96d4da", "default_location": "广州" } get_weather: { "api_key": "a861d0d5e7bf4ee1a83d9a9e4f96d4da", "default_location": "广州" }
# 获取新闻插件的配置,这里根据需要的新闻类型传入对应的url链接,默认支持社会、科技、财经新闻 # 获取新闻插件的配置,这里根据需要的新闻类型传入对应的url链接,默认支持社会、科技、财经新闻
# 更多类型的新闻列表查看 https://www.chinanews.com.cn/rss/ # 更多类型的新闻列表查看 https://www.chinanews.com.cn/rss/
get_news: get_news_from_chinanews:
default_rss_url: "https://www.chinanews.com.cn/rss/society.xml" default_rss_url: "https://www.chinanews.com.cn/rss/society.xml"
category_urls: category_urls:
society: "https://www.chinanews.com.cn/rss/society.xml" society: "https://www.chinanews.com.cn/rss/society.xml"
world: "https://www.chinanews.com.cn/rss/world.xml" world: "https://www.chinanews.com.cn/rss/world.xml"
finance: "https://www.chinanews.com.cn/rss/finance.xml" finance: "https://www.chinanews.com.cn/rss/finance.xml"
get_news_from_newsnow: {"url": "https://newsnow.busiyi.world/api/s?id="}
home_assistant: home_assistant:
devices: devices:
- 客厅,玩具灯,switch.cuco_cn_460494544_cp1_on_p_2_1 - 客厅,玩具灯,switch.cuco_cn_460494544_cp1_on_p_2_1
@@ -156,7 +157,7 @@ selected_module:
Memory: nomem Memory: nomem
# 意图识别模块开启后,可以播放音乐、控制音量、识别退出指令。 # 意图识别模块开启后,可以播放音乐、控制音量、识别退出指令。
# 不想开通意图识别,就设置成:nointent # 不想开通意图识别,就设置成:nointent
# 意图识别可使用intent_llm。优点:通用性强,缺点:增加串行前置意图识别模块,会增加处理时间,这个意图识别暂时不支持控制音量大小等iot操作 # 意图识别可使用intent_llm。优点:通用性强,缺点:增加串行前置意图识别模块,会增加处理时间,支持控制音量大小等iot操作
# 意图识别可使用function_call,缺点:需要所选择的LLM支持function_call,优点:按需调用工具、速度快,理论上能全部操作所有iot指令 # 意图识别可使用function_call,缺点:需要所选择的LLM支持function_call,优点:按需调用工具、速度快,理论上能全部操作所有iot指令
# 默认免费的ChatGLMLLM就已经支持function_call,但是如果像追求稳定建议把LLM设置成:DoubaoLLM,使用的具体model_name是:doubao-1-5-pro-32k-250115 # 默认免费的ChatGLMLLM就已经支持function_call,但是如果像追求稳定建议把LLM设置成:DoubaoLLM,使用的具体model_name是:doubao-1-5-pro-32k-250115
Intent: function_call Intent: function_call
@@ -174,6 +175,13 @@ Intent:
# 如果这里不填,则会默认使用selected_module.LLM的模型作为意图识别的思考模型 # 如果这里不填,则会默认使用selected_module.LLM的模型作为意图识别的思考模型
# 如果你的不想使用selected_module.LLM意图识别,这里最好使用独立的LLM作为意图识别,例如使用免费的ChatGLMLLM # 如果你的不想使用selected_module.LLM意图识别,这里最好使用独立的LLM作为意图识别,例如使用免费的ChatGLMLLM
llm: ChatGLMLLM llm: ChatGLMLLM
# plugins_func/functions下的模块,可以通过配置,选择加载哪个模块,加载后对话支持相应的function调用
# 系统默认已经记载“handle_exit_intent(退出识别)”、“play_music(音乐播放)”插件,请勿重复加载
# 下面是加载查天气、角色切换、加载查新闻的插件示例
functions:
- get_weather
- get_news_from_newsnow
- play_music
function_call: function_call:
# 不需要动type # 不需要动type
type: function_call type: function_call
@@ -183,7 +191,8 @@ Intent:
functions: functions:
- change_role - change_role
- get_weather - get_weather
- get_news # - get_news_from_chinanews
- get_news_from_newsnow
# play_music是服务器自带的音乐播放,hass_play_music是通过home assistant控制的独立外部程序音乐播放 # play_music是服务器自带的音乐播放,hass_play_music是通过home assistant控制的独立外部程序音乐播放
# 如果用了hass_play_music,就不要开启play_music,两者只留一个 # 如果用了hass_play_music,就不要开启play_music,两者只留一个
- play_music - play_music
@@ -221,6 +230,8 @@ ASR:
type: fun_server type: fun_server
host: 127.0.0.1 host: 127.0.0.1
port: 10096 port: 10096
is_ssl: true
output_dir: tmp/
SherpaASR: SherpaASR:
type: sherpa_onnx_local type: sherpa_onnx_local
model_dir: models/sherpa-onnx-sense-voice-zh-en-ja-ko-yue-2024-07-17 model_dir: models/sherpa-onnx-sense-voice-zh-en-ja-ko-yue-2024-07-17
@@ -253,9 +264,19 @@ ASR:
type: aliyun type: aliyun
appkey: 你的阿里云智能语音交互服务项目Appkey appkey: 你的阿里云智能语音交互服务项目Appkey
token: 你的阿里云智能语音交互服务AccessToken,临时的24小时,要长期用下方的access_key_idaccess_key_secret token: 你的阿里云智能语音交互服务AccessToken,临时的24小时,要长期用下方的access_key_idaccess_key_secret
access_key_id: 的阿里云账号access_key_id access_key_id: 的阿里云账号access_key_id
access_key_secret: 你的阿里云账号access_key_secret access_key_secret: 你的阿里云账号access_key_secret
output_dir: tmp/ output_dir: tmp/
BaiduASR:
# 获取AppID、API Key、Secret Keyhttps://console.bce.baidu.com/ai-engine/old/#/ai/speech/app/list
# 查看资源额度:https://console.bce.baidu.com/ai-engine/old/#/ai/speech/overview/resource/list
type: baidu
app_id: 你的百度语音技术AppID
api_key: 你的百度语音技术APIKey
secret_key: 你的百度语音技术SecretKey
# 语言参数,1537为普通话,具体参考:https://ai.baidu.com/ai-doc/SPEECH/0lbxfnc9b
dev_pid: 1537
output_dir: tmp/
VAD: VAD:
SileroVAD: SileroVAD:
+11 -5
View File
@@ -4,14 +4,20 @@ 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
SERVER_VERSION = "0.3.13" SERVER_VERSION = "0.4.2"
def get_module_abbreviation(module_name, module_dict): def get_module_abbreviation(module_name, module_dict):
"""获取模块名称的缩写,如果为空则返回00""" """获取模块名称的缩写,如果为空则返回00
return ( 如果名称中包含下划线,则返回下划线后面的前两个字符
module_dict.get(module_name, "")[:2] if module_dict.get(module_name) else "00" """
) module_value = module_dict.get(module_name, "")
if not module_value:
return "00"
if "_" in module_value:
parts = module_value.split("_")
return parts[-1][:2] if parts[-1] else "00"
return module_value[:2]
def build_module_string(selected_module): def build_module_string(selected_module):
@@ -146,24 +146,12 @@ def get_agent_models(
def report( def report(
mac_address: str, session_id: str, chat_type: int, content: str, opus_data mac_address: str, session_id: str, chat_type: int, content: str, audio
) -> Optional[Dict]: ) -> Optional[Dict]:
"""带熔断的业务方法示例""" """带熔断的业务方法示例"""
if not content or not ManageApiClient._instance: if not content or not ManageApiClient._instance:
return None return None
try: try:
# 处理opus_data为列表的情况
if isinstance(opus_data, list):
# 将列表中的所有bytes数据合并
combined_data = b"".join(opus_data)
else:
combined_data = opus_data
# 将二进制数据转换为Base64编码的字符串
opus_data_base64 = (
base64.b64encode(combined_data).decode("utf-8") if combined_data else None
)
return ManageApiClient._instance._execute_request( return ManageApiClient._instance._execute_request(
"POST", "POST",
f"/agent/chat-history/report", f"/agent/chat-history/report",
@@ -172,7 +160,9 @@ def report(
"sessionId": session_id, "sessionId": session_id,
"chatType": chat_type, "chatType": chat_type,
"content": content, "content": content,
"opusDataBase64": opus_data_base64, "audioBase64": (
base64.b64encode(audio).decode("utf-8") if audio else None
),
}, },
) )
except Exception as e: except Exception as e:
+130 -155
View File
@@ -1,6 +1,8 @@
import os import os
import copy import copy
import json import json
import subprocess
import sys
import uuid import uuid
import time import time
import queue import queue
@@ -18,6 +20,8 @@ from core.utils.util import (
get_string_no_punctuation_or_emoji, get_string_no_punctuation_or_emoji,
extract_json_from_string, extract_json_from_string,
initialize_modules, initialize_modules,
check_vad_update,
check_asr_update,
) )
from concurrent.futures import ThreadPoolExecutor, TimeoutError from concurrent.futures import ThreadPoolExecutor, TimeoutError
from core.handle.sendAudioHandle import sendAudioMessage from core.handle.sendAudioHandle import sendAudioMessage
@@ -52,11 +56,13 @@ class ConnectionHandler:
_intent, _intent,
server=None, server=None,
): ):
self.config = config self.common_config = config
self.server = server self.config = copy.deepcopy(config)
self.session_id = str(uuid.uuid4())
self.logger = setup_logging() self.logger = setup_logging()
self.auth = AuthMiddleware(config) self.server = server # 保存server实例的引用
self.auth = AuthMiddleware(config)
self.need_bind = False self.need_bind = False
self.bind_code = None self.bind_code = None
self.read_config_from_api = self.config.get("read_config_from_api", False) self.read_config_from_api = self.config.get("read_config_from_api", False)
@@ -66,7 +72,6 @@ class ConnectionHandler:
self.device_id = None self.device_id = None
self.client_ip = None self.client_ip = None
self.client_ip_info = {} self.client_ip_info = {}
self.session_id = None
self.prompt = None self.prompt = None
self.welcome_msg = None self.welcome_msg = None
self.max_output_size = 0 self.max_output_size = 0
@@ -87,8 +92,10 @@ class ConnectionHandler:
self.tts_report_thread = None self.tts_report_thread = None
# 依赖的组件 # 依赖的组件
self.vad = _vad self.vad = None
self.asr = _asr self.asr = None
self._asr = _asr
self._vad = _vad
self.llm = _llm self.llm = _llm
self.tts = _tts self.tts = _tts
self.memory = _memory self.memory = _memory
@@ -123,14 +130,18 @@ class ConnectionHandler:
if len(cmd) > self.max_cmd_length: if len(cmd) > self.max_cmd_length:
self.max_cmd_length = len(cmd) self.max_cmd_length = len(cmd)
self.close_after_chat = False # 是否在聊天结束后关闭连接 # 是否在聊天结束后关闭连接
self.use_function_call_mode = False self.close_after_chat = False
self.load_function_plugin = False
self.intent_type = "nointent"
self.timeout_task = None self.timeout_task = None
self.timeout_seconds = ( self.timeout_seconds = (
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"
async def handle_connection(self, ws): async def handle_connection(self, ws):
try: try:
# 获取并验证headers # 获取并验证headers
@@ -151,11 +162,9 @@ class ConnectionHandler:
self.headers["device-id"] = query_params["device-id"][0] self.headers["device-id"] = query_params["device-id"][0]
self.headers["client-id"] = query_params["client-id"][0] self.headers["client-id"] = query_params["client-id"][0]
else: else:
self.logger.bind(tag=TAG).error( await ws.send("端口正常,如需测试连接,请使用test_page.html")
"无法从请求头和URL查询参数中获取device-id" await self.close(ws)
)
return return
# 获取客户端ip地址 # 获取客户端ip地址
self.client_ip = ws.remote_address[0] self.client_ip = ws.remote_address[0]
self.logger.bind(tag=TAG).info( self.logger.bind(tag=TAG).info(
@@ -168,7 +177,6 @@ class ConnectionHandler:
# 认证通过,继续处理 # 认证通过,继续处理
self.websocket = ws self.websocket = ws
self.device_id = self.headers.get("device-id", None) self.device_id = self.headers.get("device-id", None)
self.session_id = str(uuid.uuid4())
# 启动超时检查任务 # 启动超时检查任务
self.timeout_task = asyncio.create_task(self._check_timeout()) self.timeout_task = asyncio.create_task(self._check_timeout())
@@ -178,9 +186,9 @@ class ConnectionHandler:
await self.websocket.send(json.dumps(self.welcome_msg)) await self.websocket.send(json.dumps(self.welcome_msg))
# 获取差异化配置 # 获取差异化配置
private_config = self._initialize_private_config() self._initialize_private_config()
# 异步初始化 # 异步初始化
self.executor.submit(self._initialize_components, private_config) self.executor.submit(self._initialize_components)
# tts 消化线程 # tts 消化线程
self.tts_priority_thread = threading.Thread( self.tts_priority_thread = threading.Thread(
target=self._tts_priority_thread, daemon=True target=self._tts_priority_thread, daemon=True
@@ -212,7 +220,8 @@ class ConnectionHandler:
async def _save_and_close(self, ws): async def _save_and_close(self, ws):
"""保存记忆并关闭连接""" """保存记忆并关闭连接"""
try: try:
await self.memory.save_memory(self.dialogue.dialogue) if self.memory:
await self.memory.save_memory(self.dialogue.dialogue)
except Exception as e: except Exception as e:
self.logger.bind(tag=TAG).error(f"保存记忆失败: {e}") self.logger.bind(tag=TAG).error(f"保存记忆失败: {e}")
finally: finally:
@@ -234,81 +243,58 @@ class ConnectionHandler:
elif isinstance(message, bytes): elif isinstance(message, bytes):
await handleAudioMessage(self, message) await handleAudioMessage(self, message)
async def handle_config_update(self, message): async def handle_restart(self, message):
"""处理配置更新请求""" """处理服务器重启请求"""
content = message.get("content", {}) try:
new_config = content
# 遍历所有支持的配置模块 self.logger.bind(tag=TAG).info("收到服务器重启指令,准备执行...")
updated_modules = []
for config_model in ["tts", "llm", "vad", "asr", "memory", "intent"]:
if config_model not in new_config:
continue
new_content = new_config[config_model] # 发送确认响应
old_content = self.config.get(config_model, {}) await self.websocket.send(json.dumps({
"type": "server_response",
"status": "success",
"message": "服务器重启中..."
}))
# 记录配置变更 # 异步执行重启操作
self.logger.bind(tag=TAG).info( def restart_server():
f"配置更新: {config_model} 旧值: {json.dumps(old_content, ensure_ascii=False)} " """实际执行重启的方法"""
f"新值: {json.dumps(new_content, ensure_ascii=False)}" time.sleep(1)
) self.logger.bind(tag=TAG).info("执行服务器重启...")
subprocess.Popen(
# 深度合并配置 [sys.executable, "app.py"],
if isinstance(old_content, dict) and isinstance(new_content, dict): stdin=sys.stdin,
merged = {**old_content, **new_content} stdout=sys.stdout,
self.config[config_model] = merged stderr=sys.stderr,
else: start_new_session=True
self.config[config_model] = new_content
# 标记需要重新初始化的模块
if config_model in ["llm", "tts", "asr", "vad", "intent", "memory"]:
updated_modules.append(config_model)
# 同步更新 WebSocketServer 的配置
if self.server:
async with self.server.config_lock: # 使用锁确保线程安全
for config_model in updated_modules:
self.server.config[config_model].update(new_config[config_model])
# 批量初始化模块
if updated_modules:
try:
self._initialize_components(self.config)
self.logger.bind(tag=TAG).info(
f"已重新初始化模块: {', '.join(updated_modules)}"
) )
except Exception as e: os._exit(0)
self.logger.bind(tag=TAG).error(f"模块初始化失败: {str(e)}")
await self.websocket.send(
json.dumps(
{
"type": "config_update_response",
"status": "error",
"message": f"模块初始化失败: {str(e)}",
}
)
)
return
# 返回成功响应 # 使用线程执行重启避免阻塞事件循环
await self.websocket.send( threading.Thread(target=restart_server, daemon=True).start()
json.dumps(
{
"type": "config_update_response",
"status": "success",
"message": f"已更新配置: {', '.join(updated_modules)}",
}
)
)
def _initialize_components(self, private_config): except Exception as e:
self.logger.bind(tag=TAG).error(f"重启失败: {str(e)}")
await self.websocket.send(json.dumps({
"type": "server_response",
"status": "error",
"message": f"Restart failed: {str(e)}"
}))
def _initialize_components(self):
"""初始化组件""" """初始化组件"""
if private_config is not None: if self.config.get("prompt") is not None:
self._initialize_models(private_config)
else:
self.prompt = self.config["prompt"] self.prompt = self.config["prompt"]
self.change_system_prompt(self.prompt) self.change_system_prompt(self.prompt)
self.logger.bind(tag=TAG).info(
f"初始化组件: prompt成功 {self.prompt[:50]}..."
)
"""初始化本地组件"""
if self.vad is None:
self.vad = self._vad
if self.asr is None:
self.asr = self._asr
"""加载记忆""" """加载记忆"""
self._initialize_memory() self._initialize_memory()
"""加载意图识别""" """加载意图识别"""
@@ -318,7 +304,7 @@ class ConnectionHandler:
def _init_report_threads(self): def _init_report_threads(self):
"""初始化ASR和TTS上报线程""" """初始化ASR和TTS上报线程"""
if not self.read_config_from_api: if not self.read_config_from_api or self.need_bind:
return return
if self.tts_report_thread is None or not self.tts_report_thread.is_alive(): if self.tts_report_thread is None or not self.tts_report_thread.is_alive():
self.tts_report_thread = threading.Thread( self.tts_report_thread = threading.Thread(
@@ -355,54 +341,21 @@ class ConnectionHandler:
self.logger.bind(tag=TAG).error(f"获取差异化配置失败: {e}") self.logger.bind(tag=TAG).error(f"获取差异化配置失败: {e}")
private_config = {} private_config = {}
init_tts = False init_llm, init_tts, init_memory, init_intent = (
if private_config.get("TTS", None) is not None:
init_tts = True
self.config["TTS"] = private_config["TTS"]
self.config["selected_module"]["TTS"] = private_config["selected_module"][
"TTS"
]
try:
modules = initialize_modules(
self.logger,
private_config,
False,
False,
False,
init_tts,
False,
False,
)
except Exception as e:
self.logger.bind(tag=TAG).error(f"初始化组件失败: {e}")
modules = {}
if modules.get("tts", None) is not None:
self.tts = modules["tts"]
if modules.get("prompt", None) is not None:
self.change_system_prompt(modules["prompt"])
private_config["prompt"] = None
return private_config
def _initialize_models(self, private_config):
init_vad, init_asr, init_llm, init_memory, init_intent = (
False,
False, False,
False, False,
False, False,
False, False,
) )
if private_config.get("VAD", None) is not None:
init_vad = True init_vad = check_vad_update(self.common_config, private_config)
self.config["VAD"] = private_config["VAD"] init_asr = check_asr_update(self.common_config, private_config)
self.config["selected_module"]["VAD"] = private_config["selected_module"][
"VAD" if private_config.get("TTS", None) is not None:
] init_tts = True
if private_config.get("ASR", None) is not None: self.config["TTS"] = private_config["TTS"]
init_asr = True self.config["selected_module"]["TTS"] = private_config["selected_module"][
self.config["ASR"] = private_config["ASR"] "TTS"
self.config["selected_module"]["ASR"] = private_config["selected_module"][
"ASR"
] ]
if private_config.get("LLM", None) is not None: if private_config.get("LLM", None) is not None:
init_llm = True init_llm = True
@@ -422,8 +375,11 @@ class ConnectionHandler:
self.config["selected_module"]["Intent"] = private_config[ self.config["selected_module"]["Intent"] = private_config[
"selected_module" "selected_module"
]["Intent"] ]["Intent"]
if private_config.get("prompt", None) is not None:
self.config["prompt"] = private_config["prompt"]
if private_config.get("device_max_output_size", None) is not None: if private_config.get("device_max_output_size", None) is not None:
self.max_output_size = int(private_config["device_max_output_size"]) self.max_output_size = int(private_config["device_max_output_size"])
try: try:
modules = initialize_modules( modules = initialize_modules(
self.logger, self.logger,
@@ -431,13 +387,15 @@ class ConnectionHandler:
init_vad, init_vad,
init_asr, init_asr,
init_llm, init_llm,
False, init_tts,
init_memory, init_memory,
init_intent, init_intent,
) )
except Exception as e: except Exception as e:
self.logger.bind(tag=TAG).error(f"初始化组件失败: {e}") self.logger.bind(tag=TAG).error(f"初始化组件失败: {e}")
modules = {} modules = {}
if modules.get("tts", None) is not None:
self.tts = modules["tts"]
if modules.get("vad", None) is not None: if modules.get("vad", None) is not None:
self.vad = modules["vad"] self.vad = modules["vad"]
if modules.get("asr", None) is not None: if modules.get("asr", None) is not None:
@@ -454,11 +412,11 @@ class ConnectionHandler:
self.memory.init_memory(self.device_id, self.llm) self.memory.init_memory(self.device_id, self.llm)
def _initialize_intent(self): def _initialize_intent(self):
if ( self.intent_type = self.config["Intent"][
self.config["Intent"][self.config["selected_module"]["Intent"]]["type"] self.config["selected_module"]["Intent"]
== "function_call" ]["type"]
): if self.intent_type == "function_call" or self.intent_type == "intent_llm":
self.use_function_call_mode = True self.load_function_plugin = True
"""初始化意图识别模块""" """初始化意图识别模块"""
# 获取意图识别配置 # 获取意图识别配置
intent_config = self.config["Intent"] intent_config = self.config["Intent"]
@@ -515,10 +473,12 @@ class ConnectionHandler:
processed_chars = 0 # 跟踪已处理的字符位置 processed_chars = 0 # 跟踪已处理的字符位置
try: try:
# 使用带记忆的对话 # 使用带记忆的对话
future = asyncio.run_coroutine_threadsafe( memory_str = None
self.memory.query_memory(query), self.loop if self.memory is not None:
) future = asyncio.run_coroutine_threadsafe(
memory_str = future.result() self.memory.query_memory(query), self.loop
)
memory_str = future.result()
self.logger.bind(tag=TAG).debug(f"记忆内容: {memory_str}") self.logger.bind(tag=TAG).debug(f"记忆内容: {memory_str}")
llm_responses = self.llm.response( llm_responses = self.llm.response(
@@ -565,7 +525,7 @@ class ConnectionHandler:
future = self.executor.submit( future = self.executor.submit(
self.speak_and_play, segment_text, text_index self.speak_and_play, segment_text, text_index
) )
self.tts_queue.put(future) self.tts_queue.put((future, text_index))
processed_chars += len(segment_text_raw) # 更新已处理字符位置 processed_chars += len(segment_text_raw) # 更新已处理字符位置
# 处理最后剩余的文本 # 处理最后剩余的文本
@@ -579,7 +539,7 @@ class ConnectionHandler:
future = self.executor.submit( future = self.executor.submit(
self.speak_and_play, segment_text, text_index self.speak_and_play, segment_text, text_index
) )
self.tts_queue.put(future) self.tts_queue.put((future, text_index))
self.llm_finish_task = True self.llm_finish_task = True
self.dialogue.put(Message(role="assistant", content="".join(response_message))) self.dialogue.put(Message(role="assistant", content="".join(response_message)))
@@ -606,10 +566,12 @@ class ConnectionHandler:
start_time = time.time() start_time = time.time()
# 使用带记忆的对话 # 使用带记忆的对话
future = asyncio.run_coroutine_threadsafe( memory_str = None
self.memory.query_memory(query), self.loop if self.memory is not None:
) future = asyncio.run_coroutine_threadsafe(
memory_str = future.result() self.memory.query_memory(query), self.loop
)
memory_str = future.result()
# self.logger.bind(tag=TAG).info(f"对话记录: {self.dialogue.get_llm_dialogue_with_memory(memory_str)}") # self.logger.bind(tag=TAG).info(f"对话记录: {self.dialogue.get_llm_dialogue_with_memory(memory_str)}")
@@ -695,7 +657,7 @@ class ConnectionHandler:
future = self.executor.submit( future = self.executor.submit(
self.speak_and_play, segment_text, text_index self.speak_and_play, segment_text, text_index
) )
self.tts_queue.put(future) self.tts_queue.put((future, text_index))
# 更新已处理字符位置 # 更新已处理字符位置
processed_chars += len(segment_text_raw) processed_chars += len(segment_text_raw)
@@ -754,7 +716,7 @@ class ConnectionHandler:
future = self.executor.submit( future = self.executor.submit(
self.speak_and_play, segment_text, text_index self.speak_and_play, segment_text, text_index
) )
self.tts_queue.put(future) self.tts_queue.put((future, text_index))
# 存储对话内容 # 存储对话内容
if len(response_message) > 0: if len(response_message) > 0:
@@ -816,7 +778,7 @@ class ConnectionHandler:
text = result.response text = result.response
self.recode_first_last_text(text, text_index) self.recode_first_last_text(text, text_index)
future = self.executor.submit(self.speak_and_play, text, text_index) future = self.executor.submit(self.speak_and_play, text, text_index)
self.tts_queue.put(future) self.tts_queue.put((future, text_index))
self.dialogue.put(Message(role="assistant", content=text)) self.dialogue.put(Message(role="assistant", content=text))
elif result.action == Action.REQLLM: # 调用函数后再请求llm生成回复 elif result.action == Action.REQLLM: # 调用函数后再请求llm生成回复
text = result.result text = result.result
@@ -842,14 +804,20 @@ class ConnectionHandler:
) )
self.dialogue.put( self.dialogue.put(
Message(role="tool", tool_call_id=function_id, content=text) Message(
role="tool",
tool_call_id=(
str(uuid.uuid4()) if function_id is None else function_id
),
content=text,
)
) )
self.chat_with_function_calling(text, tool_call=True) self.chat_with_function_calling(text, tool_call=True)
elif result.action == Action.NOTFOUND or result.action == Action.ERROR: elif result.action == Action.NOTFOUND or result.action == Action.ERROR:
text = result.result text = result.result
self.recode_first_last_text(text, text_index) self.recode_first_last_text(text, text_index)
future = self.executor.submit(self.speak_and_play, text, text_index) future = self.executor.submit(self.speak_and_play, text, text_index)
self.tts_queue.put(future) self.tts_queue.put((future, text_index))
self.dialogue.put(Message(role="assistant", content=text)) self.dialogue.put(Message(role="assistant", content=text))
else: else:
pass pass
@@ -859,7 +827,10 @@ class ConnectionHandler:
text = None text = None
try: try:
try: try:
future = self.tts_queue.get(timeout=1) item = self.tts_queue.get(timeout=1)
if item is None:
continue
future, text_index = item # 解包获取 Future 和 text_index
except queue.Empty: except queue.Empty:
if self.stop_event.is_set(): if self.stop_event.is_set():
break break
@@ -867,11 +838,11 @@ class ConnectionHandler:
if future is None: if future is None:
continue continue
text = None text = None
opus_datas, text_index, tts_file = [], 0, None opus_datas, tts_file = [], None
try: try:
self.logger.bind(tag=TAG).debug("正在处理TTS任务...") self.logger.bind(tag=TAG).debug("正在处理TTS任务...")
tts_timeout = int(self.config.get("tts_timeout", 10)) tts_timeout = int(self.config.get("tts_timeout", 10))
tts_file, text, text_index = future.result(timeout=tts_timeout) tts_file, text, _ = future.result(timeout=tts_timeout)
if text is None or len(text) <= 0: if text is None or len(text) <= 0:
self.logger.bind(tag=TAG).error( self.logger.bind(tag=TAG).error(
f"TTS出错:{text_index}: tts text is empty" f"TTS出错:{text_index}: tts text is empty"
@@ -885,9 +856,12 @@ class ConnectionHandler:
f"TTS生成:文件路径: {tts_file}" f"TTS生成:文件路径: {tts_file}"
) )
if os.path.exists(tts_file): if os.path.exists(tts_file):
opus_datas, _ = self.tts.audio_to_opus_data(tts_file) if self.audio_format == "pcm":
audio_datas, _ = self.tts.audio_to_pcm_data(tts_file)
else:
audio_datas, _ = self.tts.audio_to_opus_data(tts_file)
# 在这里上报TTS数据(使用文件路径) # 在这里上报TTS数据(使用文件路径)
enqueue_tts_report(self, 2, text, opus_datas) enqueue_tts_report(self, 2, text, audio_datas)
else: else:
self.logger.bind(tag=TAG).error( self.logger.bind(tag=TAG).error(
f"TTS出错:文件不存在{tts_file}" f"TTS出错:文件不存在{tts_file}"
@@ -898,7 +872,7 @@ class ConnectionHandler:
self.logger.bind(tag=TAG).error(f"TTS出错: {e}") self.logger.bind(tag=TAG).error(f"TTS出错: {e}")
if not self.client_abort: if not self.client_abort:
# 如果没有中途打断就发送语音 # 如果没有中途打断就发送语音
self.audio_play_queue.put((opus_datas, text, text_index)) self.audio_play_queue.put((audio_datas, text, text_index))
if ( if (
self.tts.delete_audio_file self.tts.delete_audio_file
and tts_file is not None and tts_file is not None
@@ -929,13 +903,13 @@ class ConnectionHandler:
text = None text = None
try: try:
try: try:
opus_datas, text, text_index = self.audio_play_queue.get(timeout=1) audio_datas, text, text_index = self.audio_play_queue.get(timeout=1)
except queue.Empty: except queue.Empty:
if self.stop_event.is_set(): if self.stop_event.is_set():
break break
continue continue
future = asyncio.run_coroutine_threadsafe( future = asyncio.run_coroutine_threadsafe(
sendAudioMessage(self, opus_datas, text, text_index), self.loop sendAudioMessage(self, audio_datas, text, text_index), self.loop
) )
future.result() future.result()
except Exception as e: except Exception as e:
@@ -1091,6 +1065,7 @@ def filter_sensitive_info(config: dict) -> dict:
"personal_access_token", "personal_access_token",
"access_token", "access_token",
"token", "token",
"secret",
"access_key_secret", "access_key_secret",
"secret_key", "secret_key",
] ]
@@ -3,15 +3,16 @@ import queue
from config.logger import setup_logging from config.logger import setup_logging
TAG = __name__ TAG = __name__
logger = setup_logging()
async def handleAbortMessage(conn): async def handleAbortMessage(conn):
logger.bind(tag=TAG).info("Abort message received") conn.logger.bind(tag=TAG).info("Abort message received")
# 设置成打断状态,会自动打断llm、tts任务 # 设置成打断状态,会自动打断llm、tts任务
conn.client_abort = True conn.client_abort = True
conn.clear_queues() conn.clear_queues()
# 打断客户端说话状态 # 打断客户端说话状态
await conn.websocket.send(json.dumps({"type": "tts", "state": "stop", "session_id": conn.session_id})) await conn.websocket.send(
json.dumps({"type": "tts", "state": "stop", "session_id": conn.session_id})
)
conn.clearSpeakStatus() conn.clearSpeakStatus()
logger.bind(tag=TAG).info("Abort message received-end") conn.logger.bind(tag=TAG).info("Abort message received-end")
@@ -4,7 +4,6 @@ from plugins_func.register import FunctionRegistry, ActionResponse, Action, Tool
from plugins_func.functions.hass_init import append_devices_to_prompt from plugins_func.functions.hass_init import append_devices_to_prompt
TAG = __name__ TAG = __name__
logger = setup_logging()
class FunctionHandler: class FunctionHandler:
@@ -40,7 +39,9 @@ class FunctionHandler:
for func in self.functions_desc: for func in self.functions_desc:
func_names.append(func["function"]["name"]) func_names.append(func["function"]["name"])
# 打印当前支持的函数列表 # 打印当前支持的函数列表
logger.bind(tag=TAG).info(f"当前支持的函数列表: {func_names}") self.conn.logger.bind(tag=TAG, session_id=self.conn.session_id).info(
f"当前支持的函数列表: {func_names}"
)
return func_names return func_names
def get_functions(self): def get_functions(self):
@@ -79,7 +80,9 @@ class FunctionHandler:
func = funcItem.func func = funcItem.func
arguments = function_call_data["arguments"] arguments = function_call_data["arguments"]
arguments = json.loads(arguments) if arguments else {} arguments = json.loads(arguments) if arguments else {}
logger.bind(tag=TAG).debug(f"调用函数: {function_name}, 参数: {arguments}") self.conn.logger.bind(tag=TAG).debug(
f"调用函数: {function_name}, 参数: {arguments}"
)
if ( if (
funcItem.type == ToolType.SYSTEM_CTL funcItem.type == ToolType.SYSTEM_CTL
or funcItem.type == ToolType.IOT_CTL or funcItem.type == ToolType.IOT_CTL
@@ -94,6 +97,6 @@ class FunctionHandler:
action=Action.NOTFOUND, result="没有找到对应的函数", response="" action=Action.NOTFOUND, result="没有找到对应的函数", response=""
) )
except Exception as e: except Exception as e:
logger.bind(tag=TAG).error(f"处理function call错误: {e}") self.conn.logger.bind(tag=TAG).error(f"处理function call错误: {e}")
return None return None
+12 -5
View File
@@ -1,5 +1,4 @@
import json import json
from config.logger import setup_logging
from core.handle.sendAudioHandle import send_stt_message from core.handle.sendAudioHandle import send_stt_message
from core.utils.util import remove_punctuation_and_length from core.utils.util import remove_punctuation_and_length
import shutil import shutil
@@ -9,7 +8,6 @@ import random
import time import time
TAG = __name__ TAG = __name__
logger = setup_logging()
WAKEUP_CONFIG = { WAKEUP_CONFIG = {
"dir": "config/assets/", "dir": "config/assets/",
@@ -21,7 +19,16 @@ WAKEUP_CONFIG = {
} }
async def handleHelloMessage(conn): async def handleHelloMessage(conn, msg_json):
"""处理hello消息"""
audio_params = msg_json.get("audio_params")
if audio_params:
format = audio_params.get("format")
conn.logger.bind(tag=TAG).info(f"客户端音频格式: {format}")
conn.audio_format = format
conn.asr.set_audio_format(format)
conn.welcome_msg["audio_params"] = audio_params
await conn.websocket.send(json.dumps(conn.welcome_msg)) await conn.websocket.send(json.dumps(conn.welcome_msg))
@@ -44,7 +51,7 @@ async def checkWakeupWords(conn, text):
if file is None: if file is None:
asyncio.create_task(wakeupWordsResponse(conn)) asyncio.create_task(wakeupWordsResponse(conn))
return False return False
opus_packets, duration = conn.tts.audio_to_opus_data(file) opus_packets, _ = conn.tts.audio_to_opus_data(file)
text_hello = WAKEUP_CONFIG["text"] text_hello = WAKEUP_CONFIG["text"]
if not text_hello: if not text_hello:
text_hello = text text_hello = text
@@ -75,7 +82,7 @@ async def wakeupWordsResponse(conn):
await asyncio.sleep(1) await asyncio.sleep(1)
wait_max_time -= 1 wait_max_time -= 1
if wait_max_time <= 0: if wait_max_time <= 0:
logger.bind(tag=TAG).error("连接对象没有llm") conn.logger.bind(tag=TAG).error("连接对象没有llm")
return return
"""唤醒词响应""" """唤醒词响应"""
@@ -5,10 +5,10 @@ from core.handle.sendAudioHandle import send_stt_message
from core.handle.helloHandle import checkWakeupWords from core.handle.helloHandle import checkWakeupWords
from core.utils.util import remove_punctuation_and_length from core.utils.util import remove_punctuation_and_length
from core.utils.dialogue import Message from core.utils.dialogue import Message
from plugins_func.register import Action
from loguru import logger from loguru import logger
TAG = __name__ TAG = __name__
logger = setup_logging()
async def handle_user_intent(conn, text): async def handle_user_intent(conn, text):
@@ -19,7 +19,7 @@ async def handle_user_intent(conn, text):
if await checkWakeupWords(conn, text): if await checkWakeupWords(conn, text):
return True return True
if conn.use_function_call_mode: if conn.intent_type == "function_call":
# 使用支持function calling的聊天方法,不再进行意图分析 # 使用支持function calling的聊天方法,不再进行意图分析
return False return False
# 使用LLM进行意图分析 # 使用LLM进行意图分析
@@ -36,7 +36,7 @@ async def check_direct_exit(conn, text):
cmd_exit = conn.cmd_exit cmd_exit = conn.cmd_exit
for cmd in cmd_exit: for cmd in cmd_exit:
if text == cmd: if text == cmd:
logger.bind(tag=TAG).info(f"识别到明确的退出命令: {text}") conn.logger.bind(tag=TAG).info(f"识别到明确的退出命令: {text}")
await send_stt_message(conn, text) await send_stt_message(conn, text)
await conn.close() await conn.close()
return True return True
@@ -46,7 +46,7 @@ async def check_direct_exit(conn, text):
async def analyze_intent_with_llm(conn, text): async def analyze_intent_with_llm(conn, text):
"""使用LLM分析用户意图""" """使用LLM分析用户意图"""
if not hasattr(conn, "intent") or not conn.intent: if not hasattr(conn, "intent") or not conn.intent:
logger.bind(tag=TAG).warning("意图识别服务未初始化") conn.logger.bind(tag=TAG).warning("意图识别服务未初始化")
return None return None
# 对话历史记录 # 对话历史记录
@@ -55,7 +55,7 @@ async def analyze_intent_with_llm(conn, text):
intent_result = await conn.intent.detect_intent(conn, dialogue.dialogue, text) intent_result = await conn.intent.detect_intent(conn, dialogue.dialogue, text)
return intent_result return intent_result
except Exception as e: except Exception as e:
logger.bind(tag=TAG).error(f"意图识别失败: {str(e)}") conn.logger.bind(tag=TAG).error(f"意图识别失败: {str(e)}")
return None return None
@@ -69,7 +69,7 @@ async def process_intent_result(conn, intent_result, original_text):
# 检查是否有function_call # 检查是否有function_call
if "function_call" in intent_data: if "function_call" in intent_data:
# 直接从意图识别获取了function_call # 直接从意图识别获取了function_call
logger.bind(tag=TAG).debug( conn.logger.bind(tag=TAG).debug(
f"检测到function_call格式的意图结果: {intent_data['function_call']['name']}" f"检测到function_call格式的意图结果: {intent_data['function_call']['name']}"
) )
function_name = intent_data["function_call"]["name"] function_name = intent_data["function_call"]["name"]
@@ -102,49 +102,51 @@ async def process_intent_result(conn, intent_result, original_text):
result = conn.func_handler.handle_llm_function_call( result = conn.func_handler.handle_llm_function_call(
conn, function_call_data conn, function_call_data
) )
if result and function_name != "play_music": logger.bind(tag=TAG).debug(f"检测到Action : {result.action}")
# 获取当前最新的文本索引
text = result.response if result:
if text is None: if result.action == Action.RESPONSE: # 直接回复前端
text = result.response
if text is not None:
speak_and_play(conn, text)
elif result.action == Action.REQLLM: # 调用函数后再请求llm生成回复
text = result.result text = result.result
if text is not None: conn.dialogue.put(Message(role="tool", content=text))
text_index = ( llm_result = conn.intent.replyResult(text, original_text)
conn.tts_last_text_index + 1 if llm_result is None:
if hasattr(conn, "tts_last_text_index") llm_result = text
else 0 speak_and_play(conn, llm_result)
) elif (
conn.recode_first_last_text(text, text_index) result.action == Action.NOTFOUND
future = conn.executor.submit( or result.action == Action.ERROR
conn.speak_and_play, text, text_index ):
) text = result.result
conn.llm_finish_task = True if text is not None:
conn.tts_queue.put(future) speak_and_play(conn, text)
conn.dialogue.put(Message(role="assistant", content=text)) elif function_name != "play_music":
# For backward compatibility with original code
# 获取当前最新的文本索引
text = result.response
if text is None:
text = result.result
if text is not None:
speak_and_play(conn, text)
# 将函数执行放在线程池中 # 将函数执行放在线程池中
conn.executor.submit(process_function_call) conn.executor.submit(process_function_call)
return True return True
return False return False
except json.JSONDecodeError as e: except json.JSONDecodeError as e:
logger.bind(tag=TAG).error(f"处理意图结果时出错: {e}") conn.logger.bind(tag=TAG).error(f"处理意图结果时出错: {e}")
return False return False
def extract_text_in_brackets(s): def speak_and_play(conn, text):
""" text_index = (
从字符串中提取中括号内的文字 conn.tts_last_text_index + 1 if hasattr(conn, "tts_last_text_index") else 0
)
:param s: 输入字符串 conn.recode_first_last_text(text, text_index)
:return: 中括号内的文字,如果不存在则返回空字符串 future = conn.executor.submit(conn.speak_and_play, text, text_index)
""" conn.llm_finish_task = True
left_bracket_index = s.find("[") conn.tts_queue.put((future, text_index))
right_bracket_index = s.find("]") conn.dialogue.put(Message(role="assistant", content=text))
if (
left_bracket_index != -1
and right_bracket_index != -1
and left_bracket_index < right_bracket_index
):
return s[left_bracket_index + 1 : right_bracket_index]
else:
return ""
+23 -20
View File
@@ -10,7 +10,6 @@ from plugins_func.register import (
) )
TAG = __name__ TAG = __name__
logger = setup_logging()
def wrap_async_function(async_func): def wrap_async_function(async_func):
@@ -21,7 +20,7 @@ def wrap_async_function(async_func):
# 获取连接对象(第一个参数) # 获取连接对象(第一个参数)
conn = args[0] conn = args[0]
if not hasattr(conn, "loop"): if not hasattr(conn, "loop"):
logger.bind(tag=TAG).error("Connection对象没有loop属性") conn.logger.bind(tag=TAG).error("Connection对象没有loop属性")
return ActionResponse( return ActionResponse(
Action.ERROR, Action.ERROR,
"Connection对象没有loop属性", "Connection对象没有loop属性",
@@ -35,7 +34,7 @@ def wrap_async_function(async_func):
# 等待结果返回 # 等待结果返回
return future.result() return future.result()
except Exception as e: except Exception as e:
logger.bind(tag=TAG).error(f"运行异步函数时出错: {e}") conn.logger.bind(tag=TAG).error(f"运行异步函数时出错: {e}")
return ActionResponse(Action.ERROR, str(e), f"执行操作时出错: {e}") return ActionResponse(Action.ERROR, str(e), f"执行操作时出错: {e}")
return wrapper return wrapper
@@ -57,7 +56,7 @@ def create_iot_function(device_name, method_name, method_info):
response_failure = "操作失败" response_failure = "操作失败"
# 打印响应参数 # 打印响应参数
logger.bind(tag=TAG).debug( conn.logger.bind(tag=TAG).debug(
f"控制函数接收到的响应参数: success='{response_success}', failure='{response_failure}'" f"控制函数接收到的响应参数: success='{response_success}', failure='{response_failure}'"
) )
@@ -86,7 +85,9 @@ def create_iot_function(device_name, method_name, method_info):
return ActionResponse(Action.RESPONSE, result, response) return ActionResponse(Action.RESPONSE, result, response)
except Exception as e: except Exception as e:
logger.bind(tag=TAG).error(f"执行{device_name}{method_name}操作失败: {e}") conn.logger.bind(tag=TAG).error(
f"执行{device_name}{method_name}操作失败: {e}"
)
# 操作失败时使用大模型提供的失败响应 # 操作失败时使用大模型提供的失败响应
response = response_failure response = response_failure
@@ -104,7 +105,7 @@ def create_iot_query_function(device_name, prop_name, prop_info):
async def iot_query_function(conn, response_success=None, response_failure=None): async def iot_query_function(conn, response_success=None, response_failure=None):
try: try:
# 打印响应参数 # 打印响应参数
logger.bind(tag=TAG).info( conn.logger.bind(tag=TAG).info(
f"查询函数接收到的响应参数: success='{response_success}', failure='{response_failure}'" f"查询函数接收到的响应参数: success='{response_success}', failure='{response_failure}'"
) )
@@ -122,7 +123,9 @@ def create_iot_query_function(device_name, prop_name, prop_info):
return ActionResponse(Action.ERROR, f"属性{prop_name}不存在", response) return ActionResponse(Action.ERROR, f"属性{prop_name}不存在", response)
except Exception as e: except Exception as e:
logger.bind(tag=TAG).error(f"查询{device_name}{prop_name}时出错: {e}") conn.logger.bind(tag=TAG).error(
f"查询{device_name}{prop_name}时出错: {e}"
)
# 查询出错时使用大模型提供的失败响应 # 查询出错时使用大模型提供的失败响应
response = response_failure response = response_failure
@@ -280,7 +283,7 @@ async def handleIotDescriptors(conn, descriptors):
await asyncio.sleep(1) await asyncio.sleep(1)
wait_max_time -= 1 wait_max_time -= 1
if wait_max_time <= 0: if wait_max_time <= 0:
logger.bind(tag=TAG).debug("连接对象没有func_handler") conn.logger.bind(tag=TAG).debug("连接对象没有func_handler")
return return
"""处理物联网描述""" """处理物联网描述"""
functions_changed = False functions_changed = False
@@ -314,7 +317,7 @@ async def handleIotDescriptors(conn, descriptors):
) )
conn.iot_descriptors[descriptor["name"]] = iot_descriptor conn.iot_descriptors[descriptor["name"]] = iot_descriptor
if conn.use_function_call_mode: if conn.load_function_plugin:
# 注册或获取设备类型 # 注册或获取设备类型
type_id = register_device_type(descriptor) type_id = register_device_type(descriptor)
device_functions = device_type_registry.get_device_functions(type_id) device_functions = device_type_registry.get_device_functions(type_id)
@@ -323,7 +326,7 @@ async def handleIotDescriptors(conn, descriptors):
if hasattr(conn, "func_handler"): if hasattr(conn, "func_handler"):
for func_name in device_functions: for func_name in device_functions:
conn.func_handler.function_registry.register_function(func_name) conn.func_handler.function_registry.register_function(func_name)
logger.bind(tag=TAG).info( conn.logger.bind(tag=TAG).info(
f"注册IOT函数到function handler: {func_name}" f"注册IOT函数到function handler: {func_name}"
) )
functions_changed = True functions_changed = True
@@ -332,8 +335,8 @@ async def handleIotDescriptors(conn, descriptors):
if functions_changed and hasattr(conn, "func_handler"): if functions_changed and hasattr(conn, "func_handler"):
conn.func_handler.upload_functions_desc() conn.func_handler.upload_functions_desc()
func_names = conn.func_handler.current_support_functions() func_names = conn.func_handler.current_support_functions()
logger.bind(tag=TAG).info(f"设备类型: {type_id}") conn.logger.bind(tag=TAG).info(f"设备类型: {type_id}")
logger.bind(tag=TAG).info( conn.logger.bind(tag=TAG).info(
f"更新function描述列表完成,当前支持的函数: {func_names}" f"更新function描述列表完成,当前支持的函数: {func_names}"
) )
@@ -347,13 +350,13 @@ async def handleIotStatus(conn, states):
for k, v in state["state"].items(): for k, v in state["state"].items():
if property_item["name"] == k: if property_item["name"] == k:
if type(v) != type(property_item["value"]): if type(v) != type(property_item["value"]):
logger.bind(tag=TAG).error( conn.logger.bind(tag=TAG).error(
f"属性{property_item['name']}的值类型不匹配" f"属性{property_item['name']}的值类型不匹配"
) )
break break
else: else:
property_item["value"] = v property_item["value"] = v
logger.bind(tag=TAG).info( conn.logger.bind(tag=TAG).info(
f"物联网状态更新: {key} , {property_item['name']} = {v}" f"物联网状态更新: {key} , {property_item['name']} = {v}"
) )
break break
@@ -367,7 +370,7 @@ async def get_iot_status(conn, name, property_name):
for property_item in value.properties: for property_item in value.properties:
if property_item["name"] == property_name: if property_item["name"] == property_name:
return property_item["value"] return property_item["value"]
logger.bind(tag=TAG).warning(f"未找到设备 {name} 的属性 {property_name}") conn.logger.bind(tag=TAG).warning(f"未找到设备 {name} 的属性 {property_name}")
return None return None
@@ -378,16 +381,16 @@ async def set_iot_status(conn, name, property_name, value):
for property_item in iot_descriptor.properties: for property_item in iot_descriptor.properties:
if property_item["name"] == property_name: if property_item["name"] == property_name:
if type(value) != type(property_item["value"]): if type(value) != type(property_item["value"]):
logger.bind(tag=TAG).error( conn.logger.bind(tag=TAG).error(
f"属性{property_item['name']}的值类型不匹配" f"属性{property_item['name']}的值类型不匹配"
) )
return return
property_item["value"] = value property_item["value"] = value
logger.bind(tag=TAG).info( conn.logger.bind(tag=TAG).info(
f"物联网状态更新: {name} , {property_name} = {value}" f"物联网状态更新: {name} , {property_name} = {value}"
) )
return return
logger.bind(tag=TAG).warning(f"未找到设备 {name} 的属性 {property_name}") conn.logger.bind(tag=TAG).warning(f"未找到设备 {name} 的属性 {property_name}")
async def send_iot_conn(conn, name, method_name, parameters): async def send_iot_conn(conn, name, method_name, parameters):
@@ -409,6 +412,6 @@ async def send_iot_conn(conn, name, method_name, parameters):
command["parameters"] = parameters command["parameters"] = parameters
send_message = json.dumps({"type": "iot", "commands": [command]}) send_message = json.dumps({"type": "iot", "commands": [command]})
await conn.websocket.send(send_message) await conn.websocket.send(send_message)
logger.bind(tag=TAG).info(f"发送物联网指令: {send_message}") conn.logger.bind(tag=TAG).info(f"发送物联网指令: {send_message}")
return return
logger.bind(tag=TAG).error(f"未找到方法{method_name}") conn.logger.bind(tag=TAG).error(f"未找到方法{method_name}")
@@ -1,4 +1,3 @@
from config.logger import setup_logging
import time import time
import copy import copy
from core.utils.util import remove_punctuation_and_length from core.utils.util import remove_punctuation_and_length
@@ -6,16 +5,18 @@ from core.handle.sendAudioHandle import send_stt_message
from core.handle.intentHandler import handle_user_intent from core.handle.intentHandler import handle_user_intent
from core.utils.output_counter import check_device_output_limit from core.utils.output_counter import check_device_output_limit
from core.handle.ttsReportHandle import enqueue_tts_report from core.handle.ttsReportHandle import enqueue_tts_report
from core.utils.util import audio_to_data
TAG = __name__ TAG = __name__
logger = setup_logging()
async def handleAudioMessage(conn, audio): async def handleAudioMessage(conn, audio):
if not conn.asr_server_receive: if conn.vad is None:
logger.bind(tag=TAG).debug(f"前期数据处理中,暂停接收")
return return
if conn.client_listen_mode == "auto": if not conn.asr_server_receive:
conn.logger.bind(tag=TAG).debug(f"前期数据处理中,暂停接收")
return
if conn.client_listen_mode == "auto" or conn.client_listen_mode == "realtime":
have_voice = conn.vad.is_vad(conn, audio) have_voice = conn.vad.is_vad(conn, audio)
else: else:
have_voice = conn.client_have_voice have_voice = conn.client_have_voice
@@ -39,7 +40,7 @@ async def handleAudioMessage(conn, audio):
conn.asr_server_receive = True conn.asr_server_receive = True
else: else:
text, _ = await conn.asr.speech_to_text(conn.asr_audio, conn.session_id) text, _ = await conn.asr.speech_to_text(conn.asr_audio, conn.session_id)
logger.bind(tag=TAG).info(f"识别文本: {text}") conn.logger.bind(tag=TAG).info(f"识别文本: {text}")
text_len, _ = remove_punctuation_and_length(text) text_len, _ = remove_punctuation_and_length(text)
if text_len > 0: if text_len > 0:
# 使用自定义模块进行上报 # 使用自定义模块进行上报
@@ -75,7 +76,7 @@ async def startToChat(conn, text):
# 意图未被处理,继续常规聊天流程 # 意图未被处理,继续常规聊天流程
await send_stt_message(conn, text) await send_stt_message(conn, text)
if conn.use_function_call_mode: if conn.intent_type == "function_call":
# 使用支持function calling的聊天方法 # 使用支持function calling的聊天方法
conn.executor.submit(conn.chat_with_function_calling, text) conn.executor.submit(conn.chat_with_function_calling, text)
else: else:
@@ -110,7 +111,7 @@ async def max_out_size(conn):
conn.tts_last_text_index = 0 conn.tts_last_text_index = 0
conn.llm_finish_task = True conn.llm_finish_task = True
file_path = "config/assets/max_output_size.wav" file_path = "config/assets/max_output_size.wav"
opus_packets, _ = conn.tts.audio_to_opus_data(file_path) opus_packets, _ = audio_to_data(file_path)
conn.audio_play_queue.put((opus_packets, text, 0)) conn.audio_play_queue.put((opus_packets, text, 0))
conn.close_after_chat = True conn.close_after_chat = True
@@ -119,7 +120,7 @@ async def check_bind_device(conn):
if conn.bind_code: if conn.bind_code:
# 确保bind_code是6位数字 # 确保bind_code是6位数字
if len(conn.bind_code) != 6: if len(conn.bind_code) != 6:
logger.bind(tag=TAG).error(f"无效的绑定码格式: {conn.bind_code}") conn.logger.bind(tag=TAG).error(f"无效的绑定码格式: {conn.bind_code}")
text = "绑定码格式错误,请检查配置。" text = "绑定码格式错误,请检查配置。"
await send_stt_message(conn, text) await send_stt_message(conn, text)
return return
@@ -132,7 +133,7 @@ async def check_bind_device(conn):
# 播放提示音 # 播放提示音
music_path = "config/assets/bind_code.wav" music_path = "config/assets/bind_code.wav"
opus_packets, _ = conn.tts.audio_to_opus_data(music_path) opus_packets, _ = audio_to_data(music_path)
conn.audio_play_queue.put((opus_packets, text, 0)) conn.audio_play_queue.put((opus_packets, text, 0))
# 逐个播放数字 # 逐个播放数字
@@ -140,10 +141,10 @@ async def check_bind_device(conn):
try: try:
digit = conn.bind_code[i] digit = conn.bind_code[i]
num_path = f"config/assets/bind_code/{digit}.wav" num_path = f"config/assets/bind_code/{digit}.wav"
num_packets, _ = conn.tts.audio_to_opus_data(num_path) num_packets, _ = audio_to_data(num_path)
conn.audio_play_queue.put((num_packets, None, i + 1)) conn.audio_play_queue.put((num_packets, None, i + 1))
except Exception as e: except Exception as e:
logger.bind(tag=TAG).error(f"播放数字音频失败: {e}") conn.logger.bind(tag=TAG).error(f"播放数字音频失败: {e}")
continue continue
else: else:
text = f"没有找到该设备的版本信息,请正确配置 OTA地址,然后重新编译固件。" text = f"没有找到该设备的版本信息,请正确配置 OTA地址,然后重新编译固件。"
@@ -152,5 +153,5 @@ async def check_bind_device(conn):
conn.tts_last_text_index = 0 conn.tts_last_text_index = 0
conn.llm_finish_task = True conn.llm_finish_task = True
music_path = "config/assets/bind_not_found.wav" music_path = "config/assets/bind_not_found.wav"
opus_packets, _ = conn.tts.audio_to_opus_data(music_path) opus_packets, _ = audio_to_data(music_path)
conn.audio_play_queue.put((opus_packets, text, 0)) conn.audio_play_queue.put((opus_packets, text, 0))
@@ -1,11 +1,9 @@
from config.logger import setup_logging
import json import json
import asyncio import asyncio
import time import time
from core.utils.util import get_string_no_punctuation_or_emoji, analyze_emotion from core.utils.util import get_string_no_punctuation_or_emoji, analyze_emotion
TAG = __name__ TAG = __name__
logger = setup_logging()
emoji_map = { emoji_map = {
"neutral": "😶", "neutral": "😶",
@@ -49,10 +47,10 @@ async def sendAudioMessage(conn, audios, text, text_index=0):
) )
if text_index == conn.tts_first_text_index: if text_index == conn.tts_first_text_index:
logger.bind(tag=TAG).info(f"发送第一段语音: {text}") conn.logger.bind(tag=TAG).info(f"发送第一段语音: {text}")
await send_tts_message(conn, "sentence_start", text) await send_tts_message(conn, "sentence_start", text)
is_first_audio = (text_index == conn.tts_first_text_index) is_first_audio = text_index == conn.tts_first_text_index
await sendAudio(conn, audios, pre_buffer=is_first_audio) await sendAudio(conn, audios, pre_buffer=is_first_audio)
await send_tts_message(conn, "sentence_end", text) await send_tts_message(conn, "sentence_end", text)
@@ -117,7 +115,7 @@ async def send_tts_message(conn, state, text=None):
stop_tts_notify_voice = conn.config.get( stop_tts_notify_voice = conn.config.get(
"stop_tts_notify_voice", "config/assets/tts_notify.mp3" "stop_tts_notify_voice", "config/assets/tts_notify.mp3"
) )
audios, duration = conn.tts.audio_to_opus_data(stop_tts_notify_voice) audios, _ = conn.tts.audio_to_opus_data(stop_tts_notify_voice)
await sendAudio(conn, audios) await sendAudio(conn, audios)
# 清除服务端讲话状态 # 清除服务端讲话状态
conn.clearSpeakStatus() conn.clearSpeakStatus()
+56 -7
View File
@@ -1,4 +1,3 @@
from config.logger import setup_logging
import json import json
from core.handle.abortHandle import handleAbortMessage from core.handle.abortHandle import handleAbortMessage
from core.handle.helloHandle import handleHelloMessage from core.handle.helloHandle import handleHelloMessage
@@ -10,25 +9,26 @@ from core.handle.ttsReportHandle import enqueue_tts_report
import asyncio import asyncio
TAG = __name__ TAG = __name__
logger = setup_logging()
async def handleTextMessage(conn, message): async def handleTextMessage(conn, message):
"""处理文本消息""" """处理文本消息"""
logger.bind(tag=TAG).info(f"收到文本消息:{message}") conn.logger.bind(tag=TAG).info(f"收到文本消息:{message}")
try: try:
msg_json = json.loads(message) msg_json = json.loads(message)
if isinstance(msg_json, int): if isinstance(msg_json, int):
await conn.websocket.send(message) await conn.websocket.send(message)
return return
if msg_json["type"] == "hello": if msg_json["type"] == "hello":
await handleHelloMessage(conn) await handleHelloMessage(conn, msg_json)
elif msg_json["type"] == "abort": elif msg_json["type"] == "abort":
await handleAbortMessage(conn) await handleAbortMessage(conn)
elif msg_json["type"] == "listen": elif msg_json["type"] == "listen":
if "mode" in msg_json: if "mode" in msg_json:
conn.client_listen_mode = msg_json["mode"] conn.client_listen_mode = msg_json["mode"]
logger.bind(tag=TAG).debug(f"客户端拾音模式:{conn.client_listen_mode}") conn.logger.bind(tag=TAG).debug(
f"客户端拾音模式:{conn.client_listen_mode}"
)
if msg_json["state"] == "start": if msg_json["state"] == "start":
conn.client_have_voice = True conn.client_have_voice = True
conn.client_voice_stop = False conn.client_voice_stop = False
@@ -80,7 +80,7 @@ async def handleTextMessage(conn, message):
await conn.websocket.send( await conn.websocket.send(
json.dumps( json.dumps(
{ {
"type": "config_update_response", "type": "server",
"status": "error", "status": "error",
"message": "服务器密钥验证失败", "message": "服务器密钥验证失败",
} }
@@ -89,6 +89,55 @@ async def handleTextMessage(conn, message):
return return
# 动态更新配置 # 动态更新配置
if msg_json["action"] == "update_config": if msg_json["action"] == "update_config":
await conn.handle_config_update(msg_json) try:
# 更新WebSocketServer的配置
if not conn.server:
await conn.websocket.send(
json.dumps(
{
"type": "config_update_response",
"status": "error",
"message": "无法获取服务器实例",
}
)
)
return
if not await conn.server.update_config():
await conn.websocket.send(
json.dumps(
{
"type": "config_update_response",
"status": "error",
"message": "更新服务器配置失败",
}
)
)
return
# 发送成功响应
await conn.websocket.send(
json.dumps(
{
"type": "config_update_response",
"status": "success",
"message": "配置更新成功",
}
)
)
except Exception as e:
conn.logger.bind(tag=TAG).error(f"更新配置失败: {str(e)}")
await conn.websocket.send(
json.dumps(
{
"type": "config_update_response",
"status": "error",
"message": f"更新配置失败: {str(e)}",
}
)
)
# 重启服务器
elif msg_json["action"] == "restart":
await conn.handle_restart(msg_json)
except json.JSONDecodeError: except json.JSONDecodeError:
await conn.websocket.send(message) await conn.websocket.send(message)
@@ -9,11 +9,11 @@ TTS上报功能已集成到ConnectionHandler类中。
具体实现请参考core/connection.py中的相关代码。 具体实现请参考core/connection.py中的相关代码。
""" """
from config.logger import setup_logging import opuslib_next
from config.manage_api_client import report from config.manage_api_client import report
TAG = __name__ TAG = __name__
logger = setup_logging()
def report_tts(conn, type, text, opus_data): def report_tts(conn, type, text, opus_data):
@@ -26,20 +26,71 @@ def report_tts(conn, type, text, opus_data):
opus_data: opus音频数据 opus_data: opus音频数据
""" """
try: try:
if opus_data:
audio_data = opus_to_wav(conn, opus_data)
else:
audio_data = None
# 执行上报 # 执行上报
report( report(
mac_address=conn.device_id, mac_address=conn.device_id,
session_id=conn.session_id, session_id=conn.session_id,
chat_type=type, chat_type=type,
content=text, content=text,
opus_data=opus_data, audio=audio_data,
) )
except Exception as e: except Exception as e:
logger.bind(tag=TAG).error(f"TTS上报失败: {e}") conn.logger.bind(tag=TAG).error(f"TTS上报失败: {e}")
def opus_to_wav(conn, opus_data):
"""将Opus数据转换为WAV格式的字节流
Args:
output_dir: 输出目录(保留参数以保持接口兼容)
opus_data: opus音频数据
Returns:
bytes: WAV格式的音频数据
"""
decoder = opuslib_next.Decoder(16000, 1) # 16kHz, 单声道
pcm_data = []
for opus_packet in opus_data:
try:
pcm_frame = decoder.decode(opus_packet, 960) # 960 samples = 60ms
pcm_data.append(pcm_frame)
except opuslib_next.OpusError as e:
conn.logger.bind(tag=TAG).error(f"Opus解码错误: {e}", exc_info=True)
if not pcm_data:
raise ValueError("没有有效的PCM数据")
# 创建WAV文件头
pcm_data_bytes = b"".join(pcm_data)
num_samples = len(pcm_data_bytes) // 2 # 16-bit samples
# WAV文件头
wav_header = bytearray()
wav_header.extend(b"RIFF") # ChunkID
wav_header.extend((36 + len(pcm_data_bytes)).to_bytes(4, "little")) # ChunkSize
wav_header.extend(b"WAVE") # Format
wav_header.extend(b"fmt ") # Subchunk1ID
wav_header.extend((16).to_bytes(4, "little")) # Subchunk1Size
wav_header.extend((1).to_bytes(2, "little")) # AudioFormat (PCM)
wav_header.extend((1).to_bytes(2, "little")) # NumChannels
wav_header.extend((16000).to_bytes(4, "little")) # SampleRate
wav_header.extend((32000).to_bytes(4, "little")) # ByteRate
wav_header.extend((2).to_bytes(2, "little")) # BlockAlign
wav_header.extend((16).to_bytes(2, "little")) # BitsPerSample
wav_header.extend(b"data") # Subchunk2ID
wav_header.extend(len(pcm_data_bytes).to_bytes(4, "little")) # Subchunk2Size
# 返回完整的WAV数据
return bytes(wav_header) + pcm_data_bytes
def enqueue_tts_report(conn, type, text, opus_data): def enqueue_tts_report(conn, type, text, opus_data):
if not conn.read_config_from_api: if not conn.read_config_from_api or conn.need_bind:
return return
"""将TTS数据加入上报队列 """将TTS数据加入上报队列
@@ -52,8 +103,8 @@ def enqueue_tts_report(conn, type, text, opus_data):
# 使用连接对象的队列,传入文本和二进制数据而非文件路径 # 使用连接对象的队列,传入文本和二进制数据而非文件路径
conn.tts_report_queue.put((type, text, opus_data)) conn.tts_report_queue.put((type, text, opus_data))
logger.bind(tag=TAG).info( conn.logger.bind(tag=TAG).debug(
f"TTS数据已加入上报队列: {conn.device_id}, 音频大小: {len(opus_data)} " f"TTS数据已加入上报队列: {conn.device_id}, 音频大小: {len(opus_data)} "
) )
except Exception as e: except Exception as e:
logger.bind(tag=TAG).error(f"加入TTS上报队列失败: {text}, {e}") conn.logger.bind(tag=TAG).error(f"加入TTS上报队列失败: {text}, {e}")
+108 -69
View File
@@ -1,86 +1,125 @@
from __future__ import annotations
from datetime import timedelta from datetime import timedelta
from typing import Optional import asyncio, os, shutil, concurrent.futures
from contextlib import AsyncExitStack from contextlib import AsyncExitStack
import os, shutil from typing import Optional, List, Dict, Any
from mcp import ClientSession, StdioServerParameters 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 config.logger import setup_logging from config.logger import setup_logging
TAG = __name__ TAG = __name__
class MCPClient: class MCPClient:
def __init__(self, config): def __init__(self, config: Dict[str, Any]):
# Initialize session and client objects
self.session: Optional[ClientSession] = None
self.exit_stack = AsyncExitStack()
self.logger = setup_logging() self.logger = setup_logging()
self.config = config self.config = config
self.tolls = []
self._worker_task: Optional[asyncio.Task] = None
self._ready_evt = asyncio.Event()
self._shutdown_evt = asyncio.Event()
self.session: Optional[ClientSession] = None
self.tools: List = []
async def initialize(self): async def initialize(self):
args = self.config.get("args", []) if self._worker_task:
return
self._worker_task = asyncio.create_task(self._worker(), name="MCPClientWorker")
await self._ready_evt.wait()
command = ( self.logger.bind(tag=TAG).info(
shutil.which("npx") f"Connected, tools = {[t.name for t in self.tools]}"
if self.config["command"] == "npx"
else self.config["command"]
) )
env={**os.environ}
if self.config.get("env"):
env.update(self.config["env"])
server_params = StdioServerParameters(
command=command,
args=args,
env=env
)
stdio_transport = await self.exit_stack.enter_async_context(stdio_client(server_params))
self.stdio, self.write = stdio_transport
time_out_delta = timedelta(seconds=15)
self.session = await self.exit_stack.enter_async_context(ClientSession(read_stream=self.stdio, write_stream=self.write, read_timeout_seconds=time_out_delta))
await self.session.initialize()
# List available tools
response = await self.session.list_tools()
tools = response.tools
self.tools = tools
self.logger.bind(tag=TAG).info(f"Connected to server with tools:{[tool.name for tool in tools]}")
def has_tool(self, tool_name):
return any(tool.name == tool_name for tool in self.tools)
def get_available_tools(self):
available_tools = [{"type": "function", "function":{
"name": tool.name,
"description": tool.description,
"parameters": tool.inputSchema
} } for tool in self.tools]
return available_tools
async def call_tool(self, tool_name: str, tool_args: dict):
self.logger.bind(tag=TAG).info(f"MCPClient Calling tool {tool_name} with args: {tool_args}")
try:
response = await self.session.call_tool(tool_name, tool_args)
except Exception as e:
self.logger.bind(tag=TAG).error(f"Error calling tool {tool_name}: {e}")
from types import SimpleNamespace
error_content = SimpleNamespace(
type='text',
text=f"Error calling tool {tool_name}: {e}"
)
error_response = SimpleNamespace(
content=[error_content],
isError=True
)
return error_response
self.logger.bind(tag=TAG).info(f"MCPClient Response from tool {tool_name}: {response}")
return response
async def cleanup(self): async def cleanup(self):
"""Clean up resources""" if not self._worker_task:
await self.exit_stack.aclose() return
self._shutdown_evt.set()
try:
await asyncio.wait_for(self._worker_task, timeout=20)
except (asyncio.TimeoutError, Exception) as e:
self.logger.bind(tag=TAG).error(f"worker shutdown err: {e}")
finally:
self._worker_task = None
def has_tool(self, name: str) -> bool:
return any(t.name == name for t in self.tools)
def get_available_tools(self):
return [
{
"type": "function",
"function": {
"name": t.name,
"description": t.description,
"parameters": t.inputSchema,
},
}
for t in self.tools
]
async def call_tool(self, name: str, args: dict):
if not self.session:
raise RuntimeError("MCPClient not initialized")
loop = self._worker_task.get_loop()
coro = self.session.call_tool(name, args)
if loop is asyncio.get_running_loop():
return await coro
fut: concurrent.futures.Future = asyncio.run_coroutine_threadsafe(coro, loop)
return await asyncio.wrap_future(fut)
async def _worker(self):
async with AsyncExitStack() as stack:
try:
# 建立 StdioClient
if "command" in self.config:
cmd = (
shutil.which("npx")
if self.config["command"] == "npx"
else self.config["command"]
)
env = {**os.environ, **self.config.get("env", {})}
params = StdioServerParameters(
command=cmd,
args=self.config.get("args", []),
env=env,
)
stdio_r, stdio_w = await stack.enter_async_context(stdio_client(params))
read_stream, write_stream = stdio_r, stdio_w
# 建立SSEClient
elif "url" in self.config:
sse_r, sse_w = await stack.enter_async_context(sse_client(self.config["url"]))
read_stream, write_stream = sse_r, sse_w
else:
raise ValueError("MCPClient config must include 'command' or 'url'")
self.session = await stack.enter_async_context(
ClientSession(
read_stream=read_stream,
write_stream=write_stream,
read_timeout_seconds=timedelta(seconds=15),
)
)
await self.session.initialize()
# 获取工具
self.tools = (await self.session.list_tools()).tools
self._ready_evt.set()
# 挂起等待关闭
await self._shutdown_evt.wait()
except Exception as e:
self.logger.bind(tag=TAG).error(f"worker error: {e}")
self._ready_evt.set()
raise
+16 -16
View File
@@ -1,9 +1,9 @@
"""MCP服务管理器""" """MCP服务管理器"""
import asyncio
import os, json import os, json
from typing import Dict, Any, List from typing import Dict, Any, List
from .MCPClient import MCPClient from .MCPClient import MCPClient
from config.logger import setup_logging
from plugins_func.register import register_function, ToolType from plugins_func.register import register_function, ToolType
from config.config_loader import get_project_dir from config.config_loader import get_project_dir
@@ -18,11 +18,10 @@ class MCPManager:
初始化MCP管理器 初始化MCP管理器
""" """
self.conn = conn self.conn = conn
self.logger = setup_logging()
self.config_path = get_project_dir() + "data/.mcp_server_settings.json" self.config_path = get_project_dir() + "data/.mcp_server_settings.json"
if os.path.exists(self.config_path) == False: if os.path.exists(self.config_path) == False:
self.config_path = "" self.config_path = ""
self.logger.bind(tag=TAG).warning( self.conn.logger.bind(tag=TAG).warning(
f"请检查mcp服务配置文件:data/.mcp_server_settings.json" f"请检查mcp服务配置文件:data/.mcp_server_settings.json"
) )
self.client: Dict[str, MCPClient] = {} self.client: Dict[str, MCPClient] = {}
@@ -41,7 +40,7 @@ class MCPManager:
config = json.load(f) config = json.load(f)
return config.get("mcpServers", {}) return config.get("mcpServers", {})
except Exception as e: except Exception as e:
self.logger.bind(tag=TAG).error( self.conn.logger.bind(tag=TAG).error(
f"Error loading MCP config from {self.config_path}: {e}" f"Error loading MCP config from {self.config_path}: {e}"
) )
return {} return {}
@@ -50,9 +49,9 @@ class MCPManager:
"""初始化所有MCP服务""" """初始化所有MCP服务"""
config = self.load_config() config = self.load_config()
for name, srv_config in config.items(): for name, srv_config in config.items():
if not srv_config.get("command"): if not srv_config.get("command") and not srv_config.get("url"):
self.logger.bind(tag=TAG).warning( self.conn.logger.bind(tag=TAG).warning(
f"Skipping server {name}: command not specified" f"Skipping server {name}: neither command nor url specified"
) )
continue continue
@@ -60,7 +59,7 @@ class MCPManager:
client = MCPClient(srv_config) client = MCPClient(srv_config)
await client.initialize() await client.initialize()
self.client[name] = client self.client[name] = client
self.logger.bind(tag=TAG).info(f"Initialized MCP client: {name}") self.conn.logger.bind(tag=TAG).info(f"Initialized MCP client: {name}")
client_tools = client.get_available_tools() client_tools = client.get_available_tools()
self.tools.extend(client_tools) self.tools.extend(client_tools)
for tool in client_tools: for tool in client_tools:
@@ -73,7 +72,7 @@ class MCPManager:
) )
except Exception as e: except Exception as e:
self.logger.bind(tag=TAG).error( self.conn.logger.bind(tag=TAG).error(
f"Failed to initialize MCP server {name}: {e}" f"Failed to initialize MCP server {name}: {e}"
) )
self.conn.func_handler.upload_functions_desc() self.conn.func_handler.upload_functions_desc()
@@ -110,7 +109,7 @@ class MCPManager:
Raises: Raises:
ValueError: 工具未找到时抛出 ValueError: 工具未找到时抛出
""" """
self.logger.bind(tag=TAG).info( self.conn.logger.bind(tag=TAG).info(
f"Executing tool {tool_name} with arguments: {arguments}" f"Executing tool {tool_name} with arguments: {arguments}"
) )
for client in self.client.values(): for client in self.client.values():
@@ -120,12 +119,13 @@ class MCPManager:
raise ValueError(f"Tool {tool_name} not found in any MCP server") raise ValueError(f"Tool {tool_name} not found in any MCP server")
async def cleanup_all(self) -> None: async def cleanup_all(self) -> None:
for name, client in self.client.items(): """依次关闭所有 MCPClient,不让异常阻断整体流程。"""
for name, client in list(self.client.items()):
try: try:
await client.cleanup() await asyncio.wait_for(client.cleanup(), timeout=20)
self.logger.bind(tag=TAG).info(f"Cleaned up MCP client: {name}") self.conn.logger.bind(tag=TAG).info(f"MCP client closed: {name}")
except Exception as e: except (asyncio.TimeoutError, Exception) as e:
self.logger.bind(tag=TAG).error( self.conn.logger.bind(tag=TAG).error(
f"Error cleaning up MCP client {name}: {e}" f"Error closing MCP client {name}: {e}"
) )
self.client.clear() self.client.clear()
@@ -20,64 +20,78 @@ from core.providers.asr.base import ASRProviderBase
TAG = __name__ TAG = __name__
logger = setup_logging() logger = setup_logging()
class AccessToken: class AccessToken:
@staticmethod @staticmethod
def _encode_text(text): def _encode_text(text):
encoded_text = parse.quote_plus(text) encoded_text = parse.quote_plus(text)
return encoded_text.replace('+', '%20').replace('*', '%2A').replace('%7E', '~') return encoded_text.replace("+", "%20").replace("*", "%2A").replace("%7E", "~")
@staticmethod @staticmethod
def _encode_dict(dic): def _encode_dict(dic):
keys = dic.keys() keys = dic.keys()
dic_sorted = [(key, dic[key]) for key in sorted(keys)] dic_sorted = [(key, dic[key]) for key in sorted(keys)]
encoded_text = parse.urlencode(dic_sorted) encoded_text = parse.urlencode(dic_sorted)
return encoded_text.replace('+', '%20').replace('*', '%2A').replace('%7E', '~') return encoded_text.replace("+", "%20").replace("*", "%2A").replace("%7E", "~")
@staticmethod @staticmethod
def create_token(access_key_id, access_key_secret): def create_token(access_key_id, access_key_secret):
parameters = {'AccessKeyId': access_key_id, parameters = {
'Action': 'CreateToken', "AccessKeyId": access_key_id,
'Format': 'JSON', "Action": "CreateToken",
'RegionId': 'cn-shanghai', "Format": "JSON",
'SignatureMethod': 'HMAC-SHA1', "RegionId": "cn-shanghai",
'SignatureNonce': str(uuid.uuid1()), "SignatureMethod": "HMAC-SHA1",
'SignatureVersion': '1.0', "SignatureNonce": str(uuid.uuid1()),
'Timestamp': time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()), "SignatureVersion": "1.0",
'Version': '2019-02-28'} "Timestamp": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()),
"Version": "2019-02-28",
}
# 构造规范化的请求字符串 # 构造规范化的请求字符串
query_string = AccessToken._encode_dict(parameters) query_string = AccessToken._encode_dict(parameters)
# print('规范化的请求字符串: %s' % query_string) # print('规范化的请求字符串: %s' % query_string)
# 构造待签名字符串 # 构造待签名字符串
string_to_sign = 'GET' + '&' + AccessToken._encode_text('/') + '&' + AccessToken._encode_text(query_string) string_to_sign = (
"GET"
+ "&"
+ AccessToken._encode_text("/")
+ "&"
+ AccessToken._encode_text(query_string)
)
# print('待签名的字符串: %s' % string_to_sign) # print('待签名的字符串: %s' % string_to_sign)
# 计算签名 # 计算签名
secreted_string = hmac.new(bytes(access_key_secret + '&', encoding='utf-8'), secreted_string = hmac.new(
bytes(string_to_sign, encoding='utf-8'), bytes(access_key_secret + "&", encoding="utf-8"),
hashlib.sha1).digest() bytes(string_to_sign, encoding="utf-8"),
hashlib.sha1,
).digest()
signature = base64.b64encode(secreted_string) signature = base64.b64encode(secreted_string)
# print('签名: %s' % signature) # print('签名: %s' % signature)
# 进行URL编码 # 进行URL编码
signature = AccessToken._encode_text(signature) signature = AccessToken._encode_text(signature)
# print('URL编码后的签名: %s' % signature) # print('URL编码后的签名: %s' % signature)
# 调用服务 # 调用服务
full_url = 'http://nls-meta.cn-shanghai.aliyuncs.com/?Signature=%s&%s' % (signature, query_string) full_url = "http://nls-meta.cn-shanghai.aliyuncs.com/?Signature=%s&%s" % (
signature,
query_string,
)
# print('url: %s' % full_url) # print('url: %s' % full_url)
# 提交HTTP GET请求 # 提交HTTP GET请求
response = requests.get(full_url) response = requests.get(full_url)
if response.ok: if response.ok:
root_obj = response.json() root_obj = response.json()
key = 'Token' key = "Token"
if key in root_obj: if key in root_obj:
token = root_obj[key]['Id'] token = root_obj[key]["Id"]
expire_time = root_obj[key]['ExpireTime'] expire_time = root_obj[key]["ExpireTime"]
return token, expire_time return token, expire_time
# print(response.text) # print(response.text)
return None, None return None, None
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__()
"""阿里云ASR初始化""" """阿里云ASR初始化"""
# 新增空值判断逻辑 # 新增空值判断逻辑
self.access_key_id = config.get("access_key_id") self.access_key_id = config.get("access_key_id")
@@ -98,39 +112,34 @@ class ASRProvider(ASRProviderBase):
# 直接使用预生成的长期token # 直接使用预生成的长期token
self.token = config.get("token") self.token = config.get("token")
self.expire_time = None self.expire_time = None
# 确保输出目录存在 # 确保输出目录存在
os.makedirs(self.output_dir, exist_ok=True) os.makedirs(self.output_dir, exist_ok=True)
def _refresh_token(self): def _refresh_token(self):
"""刷新Token并记录过期时间""" """刷新Token并记录过期时间"""
if self.access_key_id and self.access_key_secret: if self.access_key_id and self.access_key_secret:
self.token, expire_time_str = AccessToken.create_token( self.token, expire_time_str = AccessToken.create_token(
self.access_key_id, self.access_key_id, self.access_key_secret
self.access_key_secret
) )
if not expire_time_str: if not expire_time_str:
raise ValueError("无法获取有效的Token过期时间") raise ValueError("无法获取有效的Token过期时间")
try: try:
#统一转换为字符串处理 # 统一转换为字符串处理
expire_str = str(expire_time_str).strip() expire_str = str(expire_time_str).strip()
if expire_str.isdigit(): if expire_str.isdigit():
expire_time = datetime.fromtimestamp(int(expire_str)) expire_time = datetime.fromtimestamp(int(expire_str))
else: else:
expire_time = datetime.strptime( expire_time = datetime.strptime(expire_str, "%Y-%m-%dT%H:%M:%SZ")
expire_str,
"%Y-%m-%dT%H:%M:%SZ"
)
self.expire_time = expire_time.timestamp() - 60 self.expire_time = expire_time.timestamp() - 60
except Exception as e: except Exception as e:
raise ValueError(f"无效的过期时间格式: {expire_str}") from e raise ValueError(f"无效的过期时间格式: {expire_str}") from e
else: else:
self.expire_time = None self.expire_time = None
if not self.token: if not self.token:
raise ValueError("无法获取有效的访问Token") raise ValueError("无法获取有效的访问Token")
@@ -145,9 +154,12 @@ class ASRProvider(ASRProviderBase):
# f"过期时间 {datetime.fromtimestamp(self.expire_time)} | " # f"过期时间 {datetime.fromtimestamp(self.expire_time)} | "
# f"剩余 {remaining:.2f}秒") # f"剩余 {remaining:.2f}秒")
return time.time() > self.expire_time return time.time() > self.expire_time
def generate_filename(self, extension=".wav"):
return os.path.join(self.output_file, f"tts-{__name__}{datetime.now().date()}@{uuid.uuid4().hex}{extension}")
def generate_filename(self, extension=".wav"):
return os.path.join(
self.output_file,
f"tts-{__name__}{datetime.now().date()}@{uuid.uuid4().hex}{extension}",
)
def _construct_request_url(self) -> str: def _construct_request_url(self) -> str:
"""构造请求URL,包含参数""" """构造请求URL,包含参数"""
@@ -159,32 +171,17 @@ class ASRProvider(ASRProviderBase):
request += "&enable_voice_detection=false" request += "&enable_voice_detection=false"
return request return request
def decode_opus(self, opus_data: List[bytes], session_id: str) -> List[bytes]: def save_audio_to_file(self, pcm_data: List[bytes], session_id: str) -> str:
"""将Opus数据解码为PCM""" """PCM数据保存为WAV文件"""
decoder = opuslib_next.Decoder(16000, 1) # 16kHz, 单声道 module_name = __name__.split(".")[-1]
pcm_data = [] file_name = f"asr_{module_name}_{session_id}_{uuid.uuid4()}.wav"
for opus_packet in opus_data:
try:
pcm_frame = decoder.decode(opus_packet, 960) # 960 samples = 60ms
pcm_data.append(pcm_frame)
except opuslib_next.OpusError as e:
logger.bind(tag=TAG).error(f"Opus解码错误: {e}", exc_info=True)
return pcm_data
def save_audio_to_file(self, opus_data: List[bytes], session_id: str) -> str:
"""将Opus音频数据解码并保存为WAV文件"""
file_name = f"asr_{session_id}.wav"
file_path = os.path.join(self.output_dir, file_name) file_path = os.path.join(self.output_dir, file_name)
pcm_data = self.decode_opus(opus_data, session_id)
with wave.open(file_path, "wb") as wf: with wave.open(file_path, "wb") as wf:
wf.setnchannels(1) # 单声道 wf.setnchannels(1) # 单声道
wf.setsampwidth(2) # 16-bit wf.setsampwidth(2) # 16-bit
wf.setframerate(self.sample_rate) wf.setframerate(self.sample_rate)
wf.writeframes(b''.join(pcm_data)) wf.writeframes(b"".join(pcm_data))
logger.bind(tag=TAG).debug(f"音频文件已保存至: {file_path}") logger.bind(tag=TAG).debug(f"音频文件已保存至: {file_path}")
return file_path return file_path
@@ -194,22 +191,22 @@ class ASRProvider(ASRProviderBase):
try: try:
# 设置HTTP头 # 设置HTTP头
headers = { headers = {
'X-NLS-Token': self.token, "X-NLS-Token": self.token,
'Content-type': 'application/octet-stream', "Content-type": "application/octet-stream",
'Content-Length': str(len(pcm_data)) "Content-Length": str(len(pcm_data)),
} }
# 创建连接并发送请求 # 创建连接并发送请求
conn = http.client.HTTPSConnection(self.host) conn = http.client.HTTPSConnection(self.host)
request_url = self._construct_request_url() request_url = self._construct_request_url()
loop = asyncio.get_event_loop() loop = asyncio.get_event_loop()
await loop.run_in_executor(None, lambda: conn.request( await loop.run_in_executor(
method='POST', None,
url=request_url, lambda: conn.request(
body=pcm_data, method="POST", url=request_url, body=pcm_data, headers=headers
headers=headers ),
)) )
# 获取响应 # 获取响应
response = await loop.run_in_executor(None, conn.getresponse) response = await loop.run_in_executor(None, conn.getresponse)
@@ -219,16 +216,16 @@ class ASRProvider(ASRProviderBase):
# 解析响应 # 解析响应
try: try:
body_json = json.loads(body) body_json = json.loads(body)
status = body_json.get('status') status = body_json.get("status")
if status == 20000000: if status == 20000000:
result = body_json.get('result', '') result = body_json.get("result", "")
logger.bind(tag=TAG).debug(f"ASR结果: {result}") logger.bind(tag=TAG).debug(f"ASR结果: {result}")
return result return result
else: else:
logger.bind(tag=TAG).error(f"ASR失败,状态码: {status}") logger.bind(tag=TAG).error(f"ASR失败,状态码: {status}")
return None return None
except ValueError: except ValueError:
logger.bind(tag=TAG).error("响应不是JSON格式") logger.bind(tag=TAG).error("响应不是JSON格式")
return None return None
@@ -237,29 +234,37 @@ class ASRProvider(ASRProviderBase):
logger.bind(tag=TAG).error(f"ASR请求失败: {e}", exc_info=True) logger.bind(tag=TAG).error(f"ASR请求失败: {e}", exc_info=True)
return None return None
async def speech_to_text(self, opus_data: List[bytes], session_id: str) -> Tuple[Optional[str], Optional[str]]: async def speech_to_text(
self, opus_data: List[bytes], session_id: str
) -> Tuple[Optional[str], Optional[str]]:
"""将语音数据转换为文本""" """将语音数据转换为文本"""
if self._is_token_expired(): if self._is_token_expired():
logger.warning("Token已过期,正在自动刷新...") logger.warning("Token已过期,正在自动刷新...")
self._refresh_token() self._refresh_token()
file_path = None
try: try:
# 解码Opus为PCM # 解码Opus为PCM
pcm_data_list = self.decode_opus(opus_data, session_id) if self.audio_format == "pcm":
combined_pcm_data = b''.join(pcm_data_list) pcm_data = opus_data
else:
pcm_data = self.decode_opus(opus_data)
combined_pcm_data = b"".join(pcm_data)
# 判断是否保存为WAV文件
if self.delete_audio_file:
pass
else:
file_path = self.save_audio_to_file(pcm_data, session_id)
# 发送请求并获取文本 # 发送请求并获取文本
text = await self._send_request(combined_pcm_data) text = await self._send_request(combined_pcm_data)
file_path = self.save_audio_to_file(opus_data, session_id)
if self.delete_audio_file:
os.remove(file_path)
logger.bind(tag=TAG).debug(f"音频文件已删除: {file_path}")
if text: if text:
return text, None return text, file_path
return "", None
return "", file_path
except Exception as e: except Exception as e:
logger.bind(tag=TAG).error(f"语音识别失败: {e}", exc_info=True) logger.bind(tag=TAG).error(f"语音识别失败: {e}", exc_info=True)
return "", None return "", file_path
@@ -0,0 +1,106 @@
import base64
import hashlib
import hmac
import json
import time
from datetime import datetime, timezone
import os
import uuid
from typing import Optional, Tuple, List
import wave
import opuslib_next
from aip import AipSpeech
from core.providers.asr.base import ASRProviderBase
from config.logger import setup_logging
TAG = __name__
logger = setup_logging()
class ASRProvider(ASRProviderBase):
def __init__(self, config: dict, delete_audio_file: bool = True):
super().__init__()
self.app_id = config.get("app_id")
self.api_key = config.get("api_key")
self.secret_key = config.get("secret_key")
dev_pid = config.get("dev_pid", "1537")
self.dev_pid = int(dev_pid) if dev_pid else 1537
self.output_dir = config.get("output_dir")
self.delete_audio_file = delete_audio_file
self.client = AipSpeech(str(self.app_id), self.api_key, self.secret_key)
# 确保输出目录存在
os.makedirs(self.output_dir, exist_ok=True)
def save_audio_to_file(self, pcm_data: List[bytes], session_id: str) -> str:
"""PCM数据保存为WAV文件"""
module_name = __name__.split(".")[-1]
file_name = f"asr_{module_name}_{session_id}_{uuid.uuid4()}.wav"
file_path = os.path.join(self.output_dir, file_name)
with wave.open(file_path, "wb") as wf:
wf.setnchannels(1)
wf.setsampwidth(2) # 2 bytes = 16-bit
wf.setframerate(16000)
wf.writeframes(b"".join(pcm_data))
return file_path
async def speech_to_text(
self, opus_data: List[bytes], session_id: str
) -> Tuple[Optional[str], Optional[str]]:
"""将语音数据转换为文本"""
if not opus_data:
logger.bind(tag=TAG).warn("音频数据为空!")
return None, None
file_path = None
try:
# 检查配置是否已设置
if not self.app_id or not self.api_key or not self.secret_key:
logger.bind(tag=TAG).error("百度语音识别配置未设置,无法进行识别")
return None, file_path
# 将Opus音频数据解码为PCM
if self.audio_format == "pcm":
pcm_data = opus_data
else:
pcm_data = self.decode_opus(opus_data)
combined_pcm_data = b"".join(pcm_data)
# 判断是否保存为WAV文件
if self.delete_audio_file:
pass
else:
self.save_audio_to_file(pcm_data, session_id)
start_time = time.time()
# 识别本地文件
result = self.client.asr(
combined_pcm_data,
"pcm",
16000,
{
"dev_pid": str(self.dev_pid),
},
)
if result and result["err_no"] == 0:
logger.bind(tag=TAG).debug(
f"百度语音识别耗时: {time.time() - start_time:.3f}s | 结果: {result}"
)
result = result["result"][0]
return result, file_path
else:
raise Exception(
f"百度语音识别失败,错误码: {result['err_no']},错误信息: {result['err_msg']}"
)
return None, file_path
except Exception as e:
logger.bind(tag=TAG).error(f"处理音频时发生错误!{e}", exc_info=True)
return None, file_path
+27 -2
View File
@@ -1,6 +1,6 @@
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import Optional, Tuple, List from typing import Optional, Tuple, List
import opuslib_next
from config.logger import setup_logging from config.logger import setup_logging
TAG = __name__ TAG = __name__
@@ -8,12 +8,37 @@ logger = setup_logging()
class ASRProviderBase(ABC): class ASRProviderBase(ABC):
def __init__(self):
self.audio_format = "opus"
@abstractmethod @abstractmethod
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文件"""
pass pass
@abstractmethod @abstractmethod
async def speech_to_text(self, opus_data: List[bytes], session_id: str) -> Tuple[Optional[str], Optional[str]]: async def speech_to_text(
self, opus_data: List[bytes], session_id: str
) -> Tuple[Optional[str], Optional[str]]:
"""将语音数据转换为文本""" """将语音数据转换为文本"""
pass pass
def set_audio_format(self, format: str) -> None:
"""设置音频格式"""
self.audio_format = format
@staticmethod
def decode_opus(opus_data: List[bytes]) -> bytes:
"""将Opus音频数据解码为PCM数据"""
decoder = opuslib_next.Decoder(16000, 1) # 16kHz, 单声道
pcm_data = []
for opus_packet in opus_data:
try:
pcm_frame = decoder.decode(opus_packet, 960) # 960 samples = 60ms
pcm_data.append(pcm_frame)
except opuslib_next.OpusError as e:
logger.bind(tag=TAG).error(f"Opus解码错误: {e}", exc_info=True)
return pcm_data
@@ -85,11 +85,12 @@ def parse_response(res):
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__()
self.appid = config.get("appid") self.appid = 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") self.output_dir = config.get("output_dir")
self.delete_audio_file = delete_audio_file self.delete_audio_file = delete_audio_file
@@ -103,7 +104,8 @@ class ASRProvider(ASRProviderBase):
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文件"""
file_name = f"asr_{session_id}_{uuid.uuid4()}.wav" module_name = __name__.split(".")[-1]
file_name = f"asr_{module_name}_{session_id}_{uuid.uuid4()}.wav"
file_path = os.path.join(self.output_dir, file_name) file_path = os.path.join(self.output_dir, file_name)
with wave.open(file_path, "wb") as wf: with wave.open(file_path, "wb") as wf:
@@ -139,8 +141,8 @@ class ASRProvider(ASRProviderBase):
"uid": str(uuid.uuid4()), "uid": str(uuid.uuid4()),
}, },
"request": { "request": {
"reqid": reqid, "reqid": reqid,
"show_utterances": False, "show_utterances": False,
"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,
@@ -225,21 +227,6 @@ class ASRProvider(ASRProviderBase):
logger.bind(tag=TAG).error(f"ASR request failed: {e}", exc_info=True) logger.bind(tag=TAG).error(f"ASR request failed: {e}", exc_info=True)
return None return None
@staticmethod
def decode_opus(opus_data: List[bytes], session_id: str) -> List[bytes]:
decoder = opuslib_next.Decoder(16000, 1) # 16kHz, 单声道
pcm_data = []
for opus_packet in opus_data:
try:
pcm_frame = decoder.decode(opus_packet, 960) # 960 samples = 60ms
pcm_data.append(pcm_frame)
except opuslib_next.OpusError as e:
logger.bind(tag=TAG).error(f"Opus解码错误: {e}", exc_info=True)
return pcm_data
@staticmethod @staticmethod
def slice_data(data: bytes, chunk_size: int) -> (list, bool): def slice_data(data: bytes, chunk_size: int) -> (list, bool):
""" """
@@ -260,16 +247,21 @@ class ASRProvider(ASRProviderBase):
self, opus_data: List[bytes], session_id: str self, opus_data: List[bytes], session_id: str
) -> Tuple[Optional[str], Optional[str]]: ) -> Tuple[Optional[str], Optional[str]]:
"""将语音数据转换为文本""" """将语音数据转换为文本"""
file_path = None
try: try:
# 合并所有opus数据包 # 合并所有opus数据包
pcm_data = self.decode_opus(opus_data, session_id) if self.audio_format == "pcm":
pcm_data = opus_data
else:
pcm_data = self.decode_opus(opus_data)
combined_pcm_data = b"".join(pcm_data) combined_pcm_data = b"".join(pcm_data)
# 判断是否保存为WAV文件 # 判断是否保存为WAV文件
if self.delete_audio_file: if self.delete_audio_file:
pass pass
else: else:
self.save_audio_to_file(pcm_data, session_id) file_path = self.save_audio_to_file(pcm_data, session_id)
# 直接使用PCM数据 # 直接使用PCM数据
# 计算分段大小 (单声道, 16bit, 16kHz采样率) # 计算分段大小 (单声道, 16bit, 16kHz采样率)
@@ -283,9 +275,9 @@ class ASRProvider(ASRProviderBase):
logger.bind(tag=TAG).debug( logger.bind(tag=TAG).debug(
f"语音识别耗时: {time.time() - start_time:.3f}s | 结果: {text}" f"语音识别耗时: {time.time() - start_time:.3f}s | 结果: {text}"
) )
return text, None return text, file_path
return "", None return "", file_path
except Exception as e: except Exception as e:
logger.bind(tag=TAG).error(f"语音识别失败: {e}", exc_info=True) logger.bind(tag=TAG).error(f"语音识别失败: {e}", exc_info=True)
return "", None return "", file_path
@@ -6,9 +6,7 @@ import io
from config.logger import setup_logging from config.logger import setup_logging
from typing import Optional, Tuple, List from typing import Optional, Tuple, List
import uuid import uuid
import opuslib_next
from core.providers.asr.base import ASRProviderBase from core.providers.asr.base import ASRProviderBase
from funasr import AutoModel from funasr import AutoModel
from funasr.utils.postprocess_utils import rich_transcription_postprocess from funasr.utils.postprocess_utils import rich_transcription_postprocess
@@ -35,6 +33,7 @@ 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__()
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") # 修正配置键名
self.delete_audio_file = delete_audio_file self.delete_audio_file = delete_audio_file
@@ -46,25 +45,16 @@ class ASRProvider(ASRProviderBase):
model=self.model_dir, model=self.model_dir,
vad_kwargs={"max_single_segment_time": 30000}, vad_kwargs={"max_single_segment_time": 30000},
disable_update=True, disable_update=True,
hub="hf" hub="hf",
# device="cuda:0", # 启用GPU加速 # device="cuda:0", # 启用GPU加速
) )
def save_audio_to_file(self, opus_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文件"""
file_name = f"asr_{session_id}_{uuid.uuid4()}.wav" module_name = __name__.split(".")[-1]
file_name = f"asr_{module_name}_{session_id}_{uuid.uuid4()}.wav"
file_path = os.path.join(self.output_dir, file_name) file_path = os.path.join(self.output_dir, file_name)
decoder = opuslib_next.Decoder(16000, 1) # 16kHz, 单声道
pcm_data = []
for opus_packet in opus_data:
try:
pcm_frame = decoder.decode(opus_packet, 960) # 960 samples = 60ms
pcm_data.append(pcm_frame)
except opuslib_next.OpusError as e:
logger.bind(tag=TAG).error(f"Opus解码错误: {e}", exc_info=True)
with wave.open(file_path, "wb") as wf: with wave.open(file_path, "wb") as wf:
wf.setnchannels(1) wf.setnchannels(1)
wf.setsampwidth(2) # 2 bytes = 16-bit wf.setsampwidth(2) # 2 bytes = 16-bit
@@ -72,35 +62,26 @@ class ASRProvider(ASRProviderBase):
wf.writeframes(b"".join(pcm_data)) wf.writeframes(b"".join(pcm_data))
return file_path return file_path
@staticmethod
def decode_opus(opus_data: List[bytes], session_id: str) -> List[bytes]:
decoder = opuslib_next.Decoder(16000, 1) # 16kHz, 单声道 async def speech_to_text(
pcm_data = [] self, opus_data: List[bytes], session_id: str
) -> Tuple[Optional[str], Optional[str]]:
for opus_packet in opus_data:
try:
pcm_frame = decoder.decode(opus_packet, 960) # 960 samples = 60ms
pcm_data.append(pcm_frame)
except opuslib_next.OpusError as e:
logger.bind(tag=TAG).error(f"Opus解码错误: {e}", exc_info=True)
return pcm_data
async def speech_to_text(self, opus_data: List[bytes], session_id: str) -> Tuple[Optional[str], Optional[str]]:
"""语音转文本主处理逻辑""" """语音转文本主处理逻辑"""
file_path = None file_path = None
try: try:
# 合并所有opus数据包 # 合并所有opus数据包
pcm_data = self.decode_opus(opus_data, session_id) if self.audio_format == "pcm":
pcm_data = opus_data
else:
pcm_data = self.decode_opus(opus_data)
combined_pcm_data = b"".join(pcm_data) combined_pcm_data = b"".join(pcm_data)
# 判断是否保存为WAV文件 # 判断是否保存为WAV文件
if self.delete_audio_file: if self.delete_audio_file:
pass pass
else: else:
self.save_audio_to_file(pcm_data, session_id) file_path = self.save_audio_to_file(pcm_data, session_id)
# 语音识别 # 语音识别
start_time = time.time() start_time = time.time()
@@ -112,19 +93,21 @@ class ASRProvider(ASRProviderBase):
batch_size_s=60, batch_size_s=60,
) )
text = rich_transcription_postprocess(result[0]["text"]) text = rich_transcription_postprocess(result[0]["text"])
logger.bind(tag=TAG).debug(f"语音识别耗时: {time.time() - start_time:.3f}s | 结果: {text}") logger.bind(tag=TAG).debug(
f"语音识别耗时: {time.time() - start_time:.3f}s | 结果: {text}"
)
return text, file_path return text, file_path
except Exception as e: except Exception as e:
logger.bind(tag=TAG).error(f"语音识别失败: {e}", exc_info=True) logger.bind(tag=TAG).error(f"语音识别失败: {e}", exc_info=True)
return "", None return "", file_path
finally: # finally:
# 文件清理逻辑 # # 文件清理逻辑
if self.delete_audio_file and file_path and os.path.exists(file_path): # if self.delete_audio_file and file_path and os.path.exists(file_path):
try: # try:
os.remove(file_path) # os.remove(file_path)
logger.bind(tag=TAG).debug(f"已删除临时音频文件: {file_path}") # logger.bind(tag=TAG).debug(f"已删除临时音频文件: {file_path}")
except Exception as e: # except Exception as e:
logger.bind(tag=TAG).error(f"文件删除失败: {file_path} | 错误: {e}") # logger.bind(tag=TAG).error(f"文件删除失败: {file_path} | 错误: {e}")
@@ -1,57 +1,61 @@
from typing import Optional, Tuple, List from typing import Optional, Tuple, List
import opuslib_next import opuslib_next
from core.providers.asr.base import ASRProviderBase from core.providers.asr.base import ASRProviderBase
import ssl import os
import ssl
import json import json
import uuid
import wave
import websockets import websockets
from config.logger import setup_logging from config.logger import setup_logging
import asyncio import asyncio
TAG = __name__ TAG = __name__
logger = setup_logging() logger = setup_logging()
class ASRProvider(ASRProviderBase): class ASRProvider(ASRProviderBase):
def __init__(self, config: dict, delete_audio_file: bool): def __init__(self, config: dict, delete_audio_file: bool):
''' """
Initialize the ASRProvider with server configuration. Initialize the ASRProvider with server configuration.
:param config: Dictionary containing 'host', 'port', and 'is_ssl'. :param config: Dictionary containing 'host', 'port', and 'is_ssl'.
:param delete_audio_file: Boolean to indicate whether to delete audio files after processing. :param delete_audio_file: Boolean to indicate whether to delete audio files after processing.
''' """
self.host = config.get('host', 'localhost') super().__init__()
self.port = config.get('port', 10095) self.host = config.get("host", "localhost")
self.is_ssl = config.get('is_ssl', True) self.port = config.get("port", 10095)
self.is_ssl = config.get("is_ssl", True)
self.output_dir = config.get("output_dir")
self.delete_audio_file = delete_audio_file self.delete_audio_file = delete_audio_file
self.uri = f"wss://{self.host}:{self.port}" if self.is_ssl else f"ws://{self.host}:{self.port}" self.uri = (
f"wss://{self.host}:{self.port}"
if self.is_ssl
else f"ws://{self.host}:{self.port}"
)
self.ssl_context = ssl.SSLContext() if self.is_ssl else None self.ssl_context = ssl.SSLContext() if self.is_ssl else None
if self.ssl_context: if self.ssl_context:
self.ssl_context.check_hostname = False self.ssl_context.check_hostname = False
self.ssl_context.verify_mode = ssl.CERT_NONE self.ssl_context.verify_mode = ssl.CERT_NONE
def save_audio_to_file(self, opus_data: List[bytes], session_id: str) -> str: def save_audio_to_file(self, pcm_data: List[bytes], session_id: str) -> str:
"""解码Opus数据保存为WAV文件""" """PCM数据保存为WAV文件"""
pass module_name = __name__.split(".")[-1]
file_name = f"asr_{module_name}_{session_id}_{uuid.uuid4()}.wav"
file_path = os.path.join(self.output_dir, file_name)
@staticmethod with wave.open(file_path, "wb") as wf:
def decode_opus(opus_data: List[bytes]) -> bytes: wf.setnchannels(1)
"""将Opus音频数据解码为PCM数据""" wf.setsampwidth(2) # 2 bytes = 16-bit
decoder = opuslib_next.Decoder(16000, 1) # 16kHz, 单声道 wf.setframerate(16000)
pcm_data = [] wf.writeframes(b"".join(pcm_data))
for opus_packet in opus_data:
try:
pcm_frame = decoder.decode(opus_packet, 960) # 960 samples = 60ms
pcm_data.append(pcm_frame)
except opuslib_next.OpusError as e:
logger.bind(tag=TAG).error(f"Opus解码错误: {e}", exc_info=True)
return b"".join(pcm_data)
return file_path
async def _receive_responses(self, ws) -> None: async def _receive_responses(self, ws) -> None:
''' """
Asynchronous generator to receive messages from the WebSocket. Asynchronous generator to receive messages from the WebSocket.
Yields each message as it is received. Yields each message as it is received.
''' """
text = "" text = ""
while True: while True:
try: try:
@@ -64,63 +68,82 @@ class ASRProvider(ASRProviderBase):
else: else:
text += response_data.get("text", "") text += response_data.get("text", "")
except asyncio.TimeoutError: except asyncio.TimeoutError:
logger.bind(tag=TAG).error("Timeout while waiting for response from WebSocket.") logger.bind(tag=TAG).error(
"Timeout while waiting for response from WebSocket."
)
break break
except websockets.exceptions.ConnectionClosed as e: except websockets.exceptions.ConnectionClosed as e:
logger.bind(tag=TAG).error(f"WebSocket connection closed: {e}") logger.bind(tag=TAG).error(f"WebSocket connection closed: {e}")
break break
return text return text
async def _send_data(self, ws, pcm_data: bytes, session_id: str) -> tuple: async def _send_data(self, ws, pcm_data: bytes, session_id: str) -> tuple:
''' """
Internal method to handle WebSocket communication. Internal method to handle WebSocket communication.
Reuses the persistent WebSocket connection if available. Reuses the persistent WebSocket connection if available.
:param pcm_data: PCM audio data to send. :param pcm_data: PCM audio data to send.
:param session_id: Unique session identifier. :param session_id: Unique session identifier.
:return: Tuple containing recognized text and optional timestamp. :return: Tuple containing recognized text and optional timestamp.
''' """
# Send initial configuration message # Send initial configuration message
config_message = json.dumps({ config_message = json.dumps(
"mode": "offline", {
"chunk_size": [5, 10, 5], "mode": "offline",
"chunk_interval": 10, "chunk_size": [5, 10, 5],
"wav_name": session_id, "chunk_interval": 10,
"is_speaking": True, "wav_name": session_id,
"itn": False "is_speaking": True,
}) "itn": False,
}
)
await ws.send(config_message) await ws.send(config_message)
logger.bind(tag=TAG).debug(f"Sent configuration message: {config_message}") logger.bind(tag=TAG).debug(f"Sent configuration message: {config_message}")
# Send PCM data # Send PCM data
await ws.send(pcm_data) await ws.send(pcm_data)
logger.bind(tag=TAG).debug(f"Sent PCM data of length: {len(pcm_data)} bytes") logger.bind(tag=TAG).debug(f"Sent PCM data of length: {len(pcm_data)} bytes")
# Indicate end of speech # Indicate end of speech
end_message = json.dumps({"is_speaking": False}) end_message = json.dumps({"is_speaking": False})
await ws.send(end_message) await ws.send(end_message)
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) -> Tuple[Optional[str], Optional[str]]: self, opus_data: List[bytes], session_id: str
''' ) -> Tuple[Optional[str], Optional[str]]:
"""
Convert speech data to text using FunASR. Convert speech data to text using FunASR.
:param opus_data: List of Opus-encoded audio data chunks. :param opus_data: List of Opus-encoded audio data chunks.
:param session_id: Unique session identifier. :param session_id: Unique session identifier.
:return: Tuple containing recognized text and optional timestamp. :return: Tuple containing recognized text and optional timestamp.
''' """
file_path = None
if self.audio_format == "pcm":
pcm_data = opus_data
else:
pcm_data = self.decode_opus(opus_data)
combined_pcm_data = b"".join(pcm_data)
pcm_data = self.decode_opus(opus_data) # 判断是否保存为WAV文件
if self.delete_audio_file:
pass
else:
file_path = self.save_audio_to_file(pcm_data, session_id)
async with websockets.connect(self.uri, subprotocols=["binary"], ping_interval=None, ssl=self.ssl_context) as ws: async with websockets.connect(
self.uri, subprotocols=["binary"], ping_interval=None, ssl=self.ssl_context
) as ws:
try: try:
# Use asyncio to handle WebSocket communication # Use asyncio to handle WebSocket communication
send_task = asyncio.create_task(self._send_data(ws, pcm_data, session_id)) send_task = asyncio.create_task(
self._send_data(ws, combined_pcm_data, session_id)
)
receive_task = asyncio.create_task(self._receive_responses(ws)) receive_task = asyncio.create_task(self._receive_responses(ws))
# Gather tasks with error handling # Gather tasks with error handling
done, pending = await asyncio.wait( done, pending = await asyncio.wait(
[send_task, receive_task], [send_task, receive_task], return_when=asyncio.FIRST_EXCEPTION
return_when=asyncio.FIRST_EXCEPTION
) )
# Cancel any pending tasks # Cancel any pending tasks
@@ -134,11 +157,16 @@ class ASRProvider(ASRProviderBase):
# Get the result from the receive task # Get the result from the receive task
result = receive_task.result() result = receive_task.result()
return result, None # Return the recognized text and timestamp (if any) return (
result,
file_path,
) # Return the recognized text and timestamp (if any)
except websockets.exceptions.ConnectionClosed as e: except websockets.exceptions.ConnectionClosed as e:
logger.bind(tag=TAG).error(f"WebSocket connection closed: {e}") logger.bind(tag=TAG).error(f"WebSocket connection closed: {e}")
return "", None return "", file_path
except Exception as e: except Exception as e:
logger.bind(tag=TAG).error(f"Error during speech-to-text conversion: {e}", exc_info=True) logger.bind(tag=TAG).error(
return "", None f"Error during speech-to-text conversion: {e}", exc_info=True
)
return "", file_path
@@ -37,6 +37,7 @@ 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__()
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")
self.delete_audio_file = delete_audio_file self.delete_audio_file = delete_audio_file
@@ -85,7 +86,8 @@ class ASRProvider(ASRProviderBase):
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文件"""
file_name = f"asr_{session_id}_{uuid.uuid4()}.wav" module_name = __name__.split(".")[-1]
file_name = f"asr_{module_name}_{session_id}_{uuid.uuid4()}.wav"
file_path = os.path.join(self.output_dir, file_name) file_path = os.path.join(self.output_dir, file_name)
with wave.open(file_path, "wb") as wf: with wave.open(file_path, "wb") as wf:
@@ -96,21 +98,6 @@ class ASRProvider(ASRProviderBase):
return file_path return file_path
@staticmethod
def decode_opus(opus_data: List[bytes], session_id: str) -> List[bytes]:
decoder = opuslib_next.Decoder(16000, 1) # 16kHz, 单声道
pcm_data = []
for opus_packet in opus_data:
try:
pcm_frame = decoder.decode(opus_packet, 960) # 960 samples = 60ms
pcm_data.append(pcm_frame)
except opuslib_next.OpusError as e:
logger.bind(tag=TAG).error(f"Opus解码错误: {e}", exc_info=True)
return pcm_data
def read_wave(self, wave_filename: str) -> Tuple[np.ndarray, int]: def read_wave(self, wave_filename: str) -> Tuple[np.ndarray, int]:
""" """
Args: Args:
@@ -143,7 +130,10 @@ class ASRProvider(ASRProviderBase):
try: try:
# 保存音频文件 # 保存音频文件
start_time = time.time() start_time = time.time()
pcm_data = self.decode_opus(opus_data, session_id) if self.audio_format == "pcm":
pcm_data = opus_data
else:
pcm_data = self.decode_opus(opus_data)
file_path = self.save_audio_to_file(pcm_data, session_id) file_path = self.save_audio_to_file(pcm_data, session_id)
logger.bind(tag=TAG).debug( logger.bind(tag=TAG).debug(
f"音频文件保存耗时: {time.time() - start_time:.3f}s | 路径: {file_path}" f"音频文件保存耗时: {time.time() - start_time:.3f}s | 路径: {file_path}"
@@ -164,7 +154,7 @@ class ASRProvider(ASRProviderBase):
except Exception as e: except Exception as e:
logger.bind(tag=TAG).error(f"语音识别失败: {e}", exc_info=True) logger.bind(tag=TAG).error(f"语音识别失败: {e}", exc_info=True)
return "", None return "", file_path
finally: finally:
# 文件清理逻辑 # 文件清理逻辑
if self.delete_audio_file and file_path and os.path.exists(file_path): if self.delete_audio_file and file_path and os.path.exists(file_path):
@@ -17,73 +17,66 @@ from config.logger import setup_logging
TAG = __name__ TAG = __name__
logger = setup_logging() logger = setup_logging()
class ASRProvider(ASRProviderBase): class ASRProvider(ASRProviderBase):
API_URL = "https://asr.tencentcloudapi.com" API_URL = "https://asr.tencentcloudapi.com"
API_VERSION = "2019-06-14" API_VERSION = "2019-06-14"
FORMAT = "pcm" # 支持的音频格式:pcm, wav, mp3 FORMAT = "pcm" # 支持的音频格式:pcm, wav, mp3
def __init__(self, config: dict, delete_audio_file: bool = True): def __init__(self, config: dict, delete_audio_file: bool = True):
super().__init__()
self.secret_id = config.get("secret_id") self.secret_id = config.get("secret_id")
self.secret_key = config.get("secret_key") self.secret_key = config.get("secret_key")
self.output_dir = config.get("output_dir") self.output_dir = config.get("output_dir")
self.delete_audio_file = delete_audio_file
# 确保输出目录存在 # 确保输出目录存在
os.makedirs(self.output_dir, exist_ok=True) os.makedirs(self.output_dir, exist_ok=True)
def save_audio_to_file(self, opus_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文件"""
module_name = __name__.split(".")[-1]
file_name = f"tencent_asr_{session_id}_{uuid.uuid4()}.wav" file_name = f"asr_{module_name}_{session_id}_{uuid.uuid4()}.wav"
file_path = os.path.join(self.output_dir, file_name) file_path = os.path.join(self.output_dir, file_name)
with wave.open(file_path, "wb") as wf: with wave.open(file_path, "wb") as wf:
wf.setnchannels(1) wf.setnchannels(1)
wf.setsampwidth(2) # 2 bytes = 16-bit wf.setsampwidth(2) # 2 bytes = 16-bit
wf.setframerate(16000) wf.setframerate(16000)
wf.writeframes(b"".join(pcm_data)) wf.writeframes(b"".join(pcm_data))
return file_path return file_path
@staticmethod async def speech_to_text(
def decode_opus(opus_data: List[bytes]) -> bytes: self, opus_data: List[bytes], session_id: str
"""将Opus音频数据解码为PCM数据""" ) -> Tuple[Optional[str], Optional[str]]:
import opuslib_next
decoder = opuslib_next.Decoder(16000, 1) # 16kHz, 单声道
pcm_data = []
for opus_packet in opus_data:
try:
pcm_frame = decoder.decode(opus_packet, 960) # 960 samples = 60ms
pcm_data.append(pcm_frame)
except opuslib_next.OpusError as e:
logger.bind(tag=TAG).error(f"Opus解码错误: {e}", exc_info=True)
return b"".join(pcm_data)
async def speech_to_text(self, opus_data: List[bytes], session_id: str) -> Tuple[Optional[str], Optional[str]]:
"""将语音数据转换为文本""" """将语音数据转换为文本"""
if not opus_data: if not opus_data:
logger.bind(tag=TAG).warn("音频数据为空!") logger.bind(tag=TAG).warn("音频数据为空!")
return None, None return None, None
file_path = None
try: try:
# 检查配置是否已设置 # 检查配置是否已设置
if not self.secret_id or not self.secret_key: if not self.secret_id or not self.secret_key:
logger.bind(tag=TAG).error("腾讯云语音识别配置未设置,无法进行识别") logger.bind(tag=TAG).error("腾讯云语音识别配置未设置,无法进行识别")
return None, None return None, file_path
# 将Opus音频数据解码为PCM # 将Opus音频数据解码为PCM
pcm_data = self.decode_opus(opus_data) if self.audio_format == "pcm":
pcm_data = opus_data
else:
pcm_data = self.decode_opus(opus_data)
combined_pcm_data = b"".join(pcm_data)
# 判断是否保存为WAV文件 # 判断是否保存为WAV文件
if self.delete_audio_file: if self.delete_audio_file:
pass pass
else: else:
self.save_audio_to_file(pcm_data, session_id) self.save_audio_to_file(pcm_data, session_id)
# 将音频数据转换为Base64编码 # 将音频数据转换为Base64编码
base64_audio = base64.b64encode(pcm_data).decode('utf-8') base64_audio = base64.b64encode(combined_pcm_data).decode("utf-8")
# 构建请求体 # 构建请求体
request_body = self._build_request_body(base64_audio) request_body = self._build_request_body(base64_audio)
@@ -94,15 +87,17 @@ class ASRProvider(ASRProviderBase):
# 发送请求 # 发送请求
start_time = time.time() start_time = time.time()
result = self._send_request(request_body, timestamp, authorization) result = self._send_request(request_body, timestamp, authorization)
if result: if result:
logger.bind(tag=TAG).debug(f"腾讯云语音识别耗时: {time.time() - start_time:.3f}s | 结果: {result}") logger.bind(tag=TAG).debug(
f"腾讯云语音识别耗时: {time.time() - start_time:.3f}s | 结果: {result}"
return result, None )
return result, file_path
except Exception as e: except Exception as e:
logger.bind(tag=TAG).error(f"处理音频时发生错误!{e}", exc_info=True) logger.bind(tag=TAG).error(f"处理音频时发生错误!{e}", exc_info=True)
return None, None return None, file_path
def _build_request_body(self, base64_audio: str) -> str: def _build_request_body(self, base64_audio: str) -> str:
"""构建请求体""" """构建请求体"""
@@ -113,7 +108,7 @@ class ASRProvider(ASRProviderBase):
"SourceType": 1, # 音频数据来源为语音文件 "SourceType": 1, # 音频数据来源为语音文件
"VoiceFormat": self.FORMAT, # 音频格式 "VoiceFormat": self.FORMAT, # 音频格式
"Data": base64_audio, # Base64编码的音频数据 "Data": base64_audio, # Base64编码的音频数据
"DataLen": len(base64_audio) # 数据长度 "DataLen": len(base64_audio), # 数据长度
} }
return json.dumps(request_map) return json.dumps(request_map)
@@ -146,9 +141,11 @@ class ASRProvider(ASRProviderBase):
action = "SentenceRecognition" # 接口名称 action = "SentenceRecognition" # 接口名称
# 构建规范头部信息,注意顺序和格式 # 构建规范头部信息,注意顺序和格式
canonical_headers = f"content-type:{content_type.lower()}\n" + \ canonical_headers = (
f"host:{host.lower()}\n" + \ f"content-type:{content_type.lower()}\n"
f"x-tc-action:{action.lower()}\n" + f"host:{host.lower()}\n"
+ f"x-tc-action:{action.lower()}\n"
)
signed_headers = "content-type;host;x-tc-action" signed_headers = "content-type;host;x-tc-action"
@@ -156,21 +153,25 @@ class ASRProvider(ASRProviderBase):
payload_hash = self._sha256_hex(request_body) payload_hash = self._sha256_hex(request_body)
# 构建规范请求字符串 # 构建规范请求字符串
canonical_request = f"{http_request_method}\n" + \ canonical_request = (
f"{canonical_uri}\n" + \ f"{http_request_method}\n"
f"{canonical_query_string}\n" + \ + f"{canonical_uri}\n"
f"{canonical_headers}\n" + \ + f"{canonical_query_string}\n"
f"{signed_headers}\n" + \ + f"{canonical_headers}\n"
f"{payload_hash}" + f"{signed_headers}\n"
+ f"{payload_hash}"
)
# 计算规范请求的哈希值 # 计算规范请求的哈希值
hashed_canonical_request = self._sha256_hex(canonical_request) hashed_canonical_request = self._sha256_hex(canonical_request)
# 构建待签名字符串 # 构建待签名字符串
string_to_sign = f"{algorithm}\n" + \ string_to_sign = (
f"{timestamp}\n" + \ f"{algorithm}\n"
f"{credential_scope}\n" + \ + f"{timestamp}\n"
f"{hashed_canonical_request}" + f"{credential_scope}\n"
+ f"{hashed_canonical_request}"
)
# 计算签名密钥 # 计算签名密钥
secret_date = self._hmac_sha256(f"TC3{self.secret_key}", date) secret_date = self._hmac_sha256(f"TC3{self.secret_key}", date)
@@ -178,13 +179,17 @@ class ASRProvider(ASRProviderBase):
secret_signing = self._hmac_sha256(secret_service, "tc3_request") secret_signing = self._hmac_sha256(secret_service, "tc3_request")
# 计算签名 # 计算签名
signature = self._bytes_to_hex(self._hmac_sha256(secret_signing, string_to_sign)) signature = self._bytes_to_hex(
self._hmac_sha256(secret_signing, string_to_sign)
)
# 构建授权头 # 构建授权头
authorization = f"{algorithm} " + \ authorization = (
f"Credential={self.secret_id}/{credential_scope}, " + \ f"{algorithm} "
f"SignedHeaders={signed_headers}, " + \ + f"Credential={self.secret_id}/{credential_scope}, "
f"Signature={signature}" + f"SignedHeaders={signed_headers}, "
+ f"Signature={signature}"
)
return timestamp, authorization return timestamp, authorization
@@ -192,7 +197,9 @@ class ASRProvider(ASRProviderBase):
logger.bind(tag=TAG).error(f"生成认证头失败: {e}", exc_info=True) logger.bind(tag=TAG).error(f"生成认证头失败: {e}", exc_info=True)
raise RuntimeError(f"生成认证头失败: {e}") raise RuntimeError(f"生成认证头失败: {e}")
def _send_request(self, request_body: str, timestamp: str, authorization: str) -> Optional[str]: def _send_request(
self, request_body: str, timestamp: str, authorization: str
) -> Optional[str]:
"""发送请求到腾讯云API""" """发送请求到腾讯云API"""
headers = { headers = {
"Content-Type": "application/json; charset=utf-8", "Content-Type": "application/json; charset=utf-8",
@@ -201,47 +208,47 @@ class ASRProvider(ASRProviderBase):
"X-TC-Action": "SentenceRecognition", "X-TC-Action": "SentenceRecognition",
"X-TC-Version": self.API_VERSION, "X-TC-Version": self.API_VERSION,
"X-TC-Timestamp": timestamp, "X-TC-Timestamp": timestamp,
"X-TC-Region": "ap-shanghai" "X-TC-Region": "ap-shanghai",
} }
try: try:
response = requests.post(self.API_URL, headers=headers, data=request_body) response = requests.post(self.API_URL, headers=headers, data=request_body)
if not response.ok: if not response.ok:
raise IOError(f"请求失败: {response.status_code} {response.reason}") raise IOError(f"请求失败: {response.status_code} {response.reason}")
response_json = response.json() response_json = response.json()
# 检查是否有错误 # 检查是否有错误
if "Response" in response_json and "Error" in response_json["Response"]: if "Response" in response_json and "Error" in response_json["Response"]:
error = response_json["Response"]["Error"] error = response_json["Response"]["Error"]
error_code = error["Code"] error_code = error["Code"]
error_message = error["Message"] error_message = error["Message"]
raise IOError(f"API返回错误: {error_code}: {error_message}") raise IOError(f"API返回错误: {error_code}: {error_message}")
# 提取识别结果 # 提取识别结果
if "Response" in response_json and "Result" in response_json["Response"]: if "Response" in response_json and "Result" in response_json["Response"]:
return response_json["Response"]["Result"] return response_json["Response"]["Result"]
else: else:
logger.bind(tag=TAG).warn(f"响应中没有识别结果: {response_json}") logger.bind(tag=TAG).warn(f"响应中没有识别结果: {response_json}")
return "" return ""
except Exception as e: except Exception as e:
logger.bind(tag=TAG).error(f"发送请求失败: {e}", exc_info=True) logger.bind(tag=TAG).error(f"发送请求失败: {e}", exc_info=True)
return None return None
def _sha256_hex(self, data: str) -> str: def _sha256_hex(self, data: str) -> str:
"""计算字符串的SHA256哈希值""" """计算字符串的SHA256哈希值"""
digest = hashlib.sha256(data.encode('utf-8')).digest() digest = hashlib.sha256(data.encode("utf-8")).digest()
return self._bytes_to_hex(digest) return self._bytes_to_hex(digest)
def _hmac_sha256(self, key, data: str) -> bytes: def _hmac_sha256(self, key, data: str) -> bytes:
"""计算HMAC-SHA256""" """计算HMAC-SHA256"""
if isinstance(key, str): if isinstance(key, str):
key = key.encode('utf-8') key = key.encode("utf-8")
return hmac.new(key, data.encode('utf-8'), hashlib.sha256).digest() return hmac.new(key, data.encode("utf-8"), hashlib.sha256).digest()
def _bytes_to_hex(self, bytes_data: bytes) -> str: def _bytes_to_hex(self, bytes_data: bytes) -> str:
"""字节数组转十六进制字符串""" """字节数组转十六进制字符串"""
return ''.join(f"{b:02x}" for b in bytes_data) return "".join(f"{b:02x}" for b in bytes_data)
@@ -9,18 +9,6 @@ logger = setup_logging()
class IntentProviderBase(ABC): class IntentProviderBase(ABC):
def __init__(self, config): def __init__(self, config):
self.config = config self.config = config
self.intent_options = [
{
"name": "handle_exit_intent",
"desc": "结束聊天, 用户发来如再见之类的表示结束的话, 不想再进行对话的时候",
},
{
"name": "play_music",
"desc": "播放音乐, 用户希望你可以播放音乐, 只用于播放音乐的意图",
},
{"name": "get_time", "desc": "获取今天日期或者当前时间信息"},
{"name": "continue_chat", "desc": "继续聊天"},
]
def set_llm(self, llm): def set_llm(self, llm):
self.llm = llm self.llm = llm
@@ -15,57 +15,72 @@ class IntentProvider(IntentProviderBase):
def __init__(self, config): def __init__(self, config):
super().__init__(config) super().__init__(config)
self.llm = None self.llm = None
self.promot = self.get_intent_system_prompt() self.promot = ""
# 添加缓存管理 # 添加缓存管理
self.intent_cache = {} # 缓存意图识别结果 self.intent_cache = {} # 缓存意图识别结果
self.cache_expiry = 600 # 缓存有效期10分钟 self.cache_expiry = 600 # 缓存有效期10分钟
self.cache_max_size = 100 # 最多缓存100个意图 self.cache_max_size = 100 # 最多缓存100个意图
self.history_count = 4 # 默认使用最近4条对话记录
def get_intent_system_prompt(self) -> str: def get_intent_system_prompt(self, functions_list: str) -> str:
""" """
根据配置的意图选项动态生成系统提示词 根据配置的意图选项和可用函数动态生成系统提示词
Args:
functions: 可用的函数列表,JSON格式字符串
Returns: Returns:
格式化后的系统提示词 格式化后的系统提示词
""" """
# 构建函数说明部分
functions_desc = "可用的函数列表:\n"
for func in functions_list:
func_info = func.get("function", {})
name = func_info.get("name", "")
desc = func_info.get("description", "")
params = func_info.get("parameters", {})
functions_desc += f"\n函数名: {name}\n"
functions_desc += f"描述: {desc}\n"
if params:
functions_desc += "参数:\n"
for param_name, param_info in params.get("properties", {}).items():
param_desc = param_info.get("description", "")
param_type = param_info.get("type", "")
functions_desc += f"- {param_name} ({param_type}): {param_desc}\n"
functions_desc += "---\n"
prompt = ( prompt = (
"你是一个意图识别助手。请分析用户的最后一句话,判断用户意图属于以下哪一类:\n" "你是一个意图识别助手。请分析用户的最后一句话,判断用户意图并调用相应的函数。\n\n"
"<start>" f"{functions_desc}\n"
f"{str(self.intent_options)}" "处理步骤:\n"
"<end>\n" "1. 分析用户输入,确定用户意图\n"
"处理步骤:" "2. 从可用函数列表中选择最匹配的函数\n"
"1. 思考意图类型,生成function_call格式" "3. 如果找到匹配的函数,生成对应的function_call 格式\n"
"\n\n" '4. 如果没有找到匹配的函数,返回{"function_call": {"name": "continue_chat"}}\n\n'
"返回格式示例\n" "返回格式要求\n"
'1. 播放音乐意图: {"function_call": {"name": "play_music", "arguments": {"song_name": "音乐名称"}}}\n' "1. 必须返回纯JSON格式\n"
'2. 结束对话意图: {"function_call": {"name": "handle_exit_intent", "arguments": {"say_goodbye": "goodbye"}}}\n' "2. 必须包含function_call字段\n"
'3. 获取当天日期时间: {"function_call": {"name": "get_time"}}\n' "3. function_call必须包含name字段\n"
'4. 继续聊天意图: {"function_call": {"name": "continue_chat"}}\n' "4. 如果函数需要参数,必须包含arguments字段\n\n"
"\n" "示例:\n"
"注意:\n"
'- 播放音乐:无歌名时,song_name设为"random"\n'
"- 如果没有明显的意图,应按照继续聊天意图处理\n"
"- 只返回纯JSON,不要任何其他内容\n"
"\n"
"示例分析:\n"
"```\n" "```\n"
"用户: 你也太搞笑了\n" "用户: 现在几点了?\n"
'返回: {"function_call": {"name": "continue_chat"}}\n'
"```\n"
"```\n"
"用户: 现在是几号了?现在几点了?\n"
'返回: {"function_call": {"name": "get_time"}}\n' '返回: {"function_call": {"name": "get_time"}}\n'
"```\n" "```\n"
"```\n" "```\n"
"用户: 我们明天再聊吧\n" "用户: 我想结束对话\n"
'返回: {"function_call": {"name": "handle_exit_intent"}}\n' '返回: {"function_call": {"name": "handle_exit_intent", "arguments": {"say_goodbye": "goodbye"}}}\n'
"```\n" "```\n"
"```\n" "```\n"
"用户: 播放中秋月\n" "用户: 你好啊\n"
'返回: {"function_call": {"name": "play_music", "arguments": {"song_name": "中秋月"}}}\n' '返回: {"function_call": {"name": "continue_chat"}}\n'
"```\n" "```\n\n"
"```\n" "注意:\n"
"可用的音乐名称:\n" "1. 只返回JSON格式,不要包含任何其他文字\n"
'2. 如果没有找到匹配的函数,返回{"function_call": {"name": "continue_chat"}}\n'
"3. 确保返回的JSON格式正确,包含所有必要的字段\n"
) )
return prompt return prompt
@@ -90,6 +105,14 @@ class IntentProvider(IntentProviderBase):
for key, _ in sorted_items[: len(sorted_items) - self.cache_max_size]: for key, _ in sorted_items[: len(sorted_items) - self.cache_max_size]:
del self.intent_cache[key] del self.intent_cache[key]
def replyResult(self, text: str, original_text: str):
llm_result = self.llm.response_no_stream(
system_prompt=text,
user_prompt="请根据以上内容,像人类一样说话的口吻回复用户,要求简洁,请直接返回结果。用户现在说:"
+ original_text,
)
return llm_result
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")
@@ -118,22 +141,35 @@ class IntentProvider(IntentProviderBase):
# 清理缓存 # 清理缓存
self.clean_cache() self.clean_cache()
# 构建用户最后一句话的提示 if self.promot == "":
msgStr = "" if hasattr(conn, "func_handler"):
functions = conn.func_handler.get_functions()
self.promot = self.get_intent_system_prompt(functions)
# 只使用最后两句即可
if len(dialogue_history) >= 2:
# 保证最少有两句话的时候处理
msgStr += f"{dialogue_history[-2].role}: {dialogue_history[-2].content}\n"
msgStr += f"{dialogue_history[-1].role}: {dialogue_history[-1].content}\n"
msgStr += f"User: {text}\n"
user_prompt = f"当前的对话如下:\n{msgStr}"
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"]
prompt_music = f"{self.promot}\n<start>{music_file_names}\n<end>" prompt_music = f"{self.promot}\n<musicNames>{music_file_names}\n</musicNames>"
devices = conn.config["plugins"]["home_assistant"].get("devices", [])
if len(devices) > 0:
hass_prompt = "\n下面是我家智能设备列表(位置,设备名,entity_id),可以通过homeassistant控制\n"
for device in devices:
hass_prompt += device + "\n"
prompt_music += hass_prompt
logger.bind(tag=TAG).debug(f"User prompt: {prompt_music}") logger.bind(tag=TAG).debug(f"User prompt: {prompt_music}")
# 构建用户对话历史的提示
msgStr = ""
# 获取最近的对话历史
start_idx = max(0, len(dialogue_history) - self.history_count)
for i in range(start_idx, len(dialogue_history)):
msgStr += f"{dialogue_history[i].role}: {dialogue_history[i].content}\n"
msgStr += f"User: {text}\n"
user_prompt = f"current dialogue:\n{msgStr}"
# 记录预处理完成时间 # 记录预处理完成时间
preprocess_time = time.time() - total_start_time preprocess_time = time.time() - total_start_time
logger.bind(tag=TAG).debug(f"意图识别预处理耗时: {preprocess_time:.4f}") logger.bind(tag=TAG).debug(f"意图识别预处理耗时: {preprocess_time:.4f}")
+9 -47
View File
@@ -1,11 +1,9 @@
import asyncio import asyncio
from config.logger import setup_logging from config.logger import setup_logging
import os import os
import numpy as np
import opuslib_next
from pydub import AudioSegment
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from core.utils.tts import MarkdownCleaner from core.utils.tts import MarkdownCleaner
from core.utils.util import audio_to_data
TAG = __name__ TAG = __name__
logger = setup_logging() logger = setup_logging()
@@ -29,7 +27,9 @@ class TTSProviderBase(ABC):
try: try:
asyncio.run(self.text_to_speak(text, tmp_file)) asyncio.run(self.text_to_speak(text, tmp_file))
except Exception as e: except Exception as e:
logger.bind(tag=TAG).warning(f"语音生成失败{5 - max_repeat_time + 1}次: {text},错误: {e}") logger.bind(tag=TAG).warning(
f"语音生成失败{5 - max_repeat_time + 1}次: {text},错误: {e}"
)
# 未执行成功,删除文件 # 未执行成功,删除文件
if os.path.exists(tmp_file): if os.path.exists(tmp_file):
os.remove(tmp_file) os.remove(tmp_file)
@@ -53,48 +53,10 @@ class TTSProviderBase(ABC):
async def text_to_speak(self, text, output_file): async def text_to_speak(self, text, output_file):
pass pass
def audio_to_pcm_data(self, audio_file_path):
"""音频文件转换为PCM编码"""
return audio_to_data(audio_file_path, is_opus=False)
def audio_to_opus_data(self, audio_file_path): def audio_to_opus_data(self, audio_file_path):
"""音频文件转换为Opus编码""" """音频文件转换为Opus编码"""
# 获取文件后缀名 return audio_to_data(audio_file_path, is_opus=True)
file_type = os.path.splitext(audio_file_path)[1]
if file_type:
file_type = file_type.lstrip(".")
# 读取音频文件,-nostdin 参数:不要从标准输入读取数据,否则FFmpeg会阻塞
audio = AudioSegment.from_file(
audio_file_path, format=file_type, parameters=["-nostdin"]
)
# 转换为单声道/16kHz采样率/16位小端编码(确保与编码器匹配)
audio = audio.set_channels(1).set_frame_rate(16000).set_sample_width(2)
# 音频时长(秒)
duration = len(audio) / 1000.0
# 获取原始PCM数据(16位小端)
raw_data = audio.raw_data
# 初始化Opus编码器
encoder = opuslib_next.Encoder(16000, 1, opuslib_next.APPLICATION_AUDIO)
# 编码参数
frame_duration = 60 # 60ms per frame
frame_size = int(16000 * frame_duration / 1000) # 960 samples/frame
opus_datas = []
# 按帧处理所有音频数据(包括最后一帧可能补零)
for i in range(0, len(raw_data), frame_size * 2): # 16bit=2bytes/sample
# 获取当前帧的二进制数据
chunk = raw_data[i : i + frame_size * 2]
# 如果最后一帧不足,补零
if len(chunk) < frame_size * 2:
chunk += b"\x00" * (frame_size * 2 - len(chunk))
# 转换为numpy数组处理
np_frame = np.frombuffer(chunk, dtype=np.int16)
# 编码Opus数据
opus_data = encoder.encode(np_frame.tobytes(), frame_size)
opus_datas.append(opus_data)
return opus_datas, duration
@@ -38,9 +38,13 @@ class TTSProvider(TTSProviderBase):
"Authorization": f"Bearer {self.access_token}", "Authorization": f"Bearer {self.access_token}",
"Content-Type": "application/json", "Content-Type": "application/json",
} }
response = requests.request(
"POST", self.api_url, json=request_json, headers=headers try:
) response = requests.request(
data = response.content "POST", self.api_url, json=request_json, headers=headers
file_to_save = open(output_file, "wb") )
file_to_save.write(data) data = response.content
file_to_save = open(output_file, "wb")
file_to_save.write(data)
except Exception as e:
raise Exception(f"{__name__} error: {e}")
@@ -32,4 +32,6 @@ class TTSProvider(TTSProviderBase):
with open(output_file, "wb") as file: with open(output_file, "wb") as file:
file.write(resp.content) file.write(resp.content)
else: else:
logger.bind(tag=TAG).error(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)
raise Exception(error_msg) # 抛出异常,让调用方捕获
@@ -51,7 +51,7 @@ class TTSProvider(TTSProviderBase):
request_json = { request_json = {
"app": { "app": {
"appid": f"{self.appid}", "appid": f"{self.appid}",
"token": "access_token", "token": self.access_token,
"cluster": self.cluster, "cluster": self.cluster,
}, },
"user": {"uid": "1"}, "user": {"uid": "1"},
+14 -10
View File
@@ -20,14 +20,18 @@ class TTSProvider(TTSProviderBase):
) )
async def text_to_speak(self, text, output_file): async def text_to_speak(self, text, output_file):
communicate = edge_tts.Communicate(text, voice=self.voice) try:
# 确保目录存在并创建空文件 communicate = edge_tts.Communicate(text, voice=self.voice)
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(): async for chunk in communicate.stream():
if chunk["type"] == "audio": # 只处理音频数据块 if chunk["type"] == "audio": # 只处理音频数据块
f.write(chunk["data"]) f.write(chunk["data"])
except Exception as e:
error_msg = f"Edge TTS请求失败: {e}"
raise Exception(error_msg) # 抛出异常,让调用方捕获
@@ -177,5 +177,7 @@ class TTSProvider(TTSProviderBase):
audio_file.write(audio_content) audio_file.write(audio_content)
else: else:
print(f"Request failed with status code {response.status_code}") error_msg = f"Request failed with status code {response.status_code}"
print(error_msg)
print(response.json()) print(response.json())
raise Exception(error_msg)
@@ -105,6 +105,6 @@ class TTSProvider(TTSProviderBase):
with open(output_file, "wb") as file: with open(output_file, "wb") as file:
file.write(resp.content) file.write(resp.content)
else: else:
logger.bind(tag=TAG).error( error_msg = f"GPT_SoVITS_V2 TTS请求失败: {resp.status_code} - {resp.text}"
f"GPT_SoVITS_V2 TTS请求失败: {resp.status_code} - {resp.text}" logger.bind(tag=TAG).error(error_msg)
) raise Exception(error_msg)
@@ -64,6 +64,7 @@ class TTSProvider(TTSProviderBase):
with open(output_file, "wb") as file: with open(output_file, "wb") as file:
file.write(resp.content) file.write(resp.content)
else: else:
logger.bind(tag=TAG).error( error_msg = f"GPT_SoVITS_V3 TTS请求失败: {resp.status_code} - {resp.text}"
f"GPT_SoVITS_V3 TTS请求失败: {resp.status_code} - {resp.text}" logger.bind(tag=TAG).error(error_msg)
) raise Exception(error_msg)
@@ -39,9 +39,12 @@ class TTSProvider(TTSProviderBase):
"Authorization": f"Bearer {self.access_token}", "Authorization": f"Bearer {self.access_token}",
"Content-Type": "application/json", "Content-Type": "application/json",
} }
response = requests.request( try:
"POST", self.api_url, json=request_json, headers=headers response = requests.request(
) "POST", self.api_url, json=request_json, headers=headers
data = response.content )
file_to_save = open(output_file, "wb") data = response.content
file_to_save.write(data) file_to_save = open(output_file, "wb")
file_to_save.write(data)
except Exception as e:
raise Exception(f"{__name__} error: {e}")
+12 -10
View File
@@ -58,8 +58,8 @@ class TTSProvider(TTSProviderBase):
resp = requests.request("POST", url, data=payload) resp = requests.request("POST", url, data=payload)
if resp.status_code != 200: if resp.status_code != 200:
logger.bind(tag=TAG).error(f"TTS请求失败: {resp.text}") logger.bind(tag=TAG).error(f"TTSON 请求失败: {resp.text}")
return None raise Exception(f"{__name__}: TTS请求失败")
resp_json = resp.json() resp_json = resp.json()
try: try:
result = ( result = (
@@ -71,13 +71,15 @@ class TTSProvider(TTSProviderBase):
+ "&voice_audio_path=" + "&voice_audio_path="
+ resp_json["voice_path"] + resp_json["voice_path"]
) )
audio_content = requests.get(result)
with open(output_file, "wb") as f:
f.write(audio_content.content)
return True
voice_path = resp_json.get("voice_path")
des_path = output_file
shutil.move(voice_path, des_path)
except Exception as e: except Exception as e:
print("error:", e) print("error:", e)
raise Exception(f"{__name__}: TTS请求失败")
audio_content = requests.get(result)
with open(output_file, "wb") as f:
f.write(audio_content.content)
return True
voice_path = resp_json.get("voice_path")
des_path = output_file
shutil.move(voice_path, des_path)
+24 -9
View File
@@ -4,7 +4,14 @@ from datetime import datetime
class Message: class Message:
def __init__(self, role: str, content: str = None, uniq_id: str = None, tool_calls = None, tool_call_id=None): def __init__(
self,
role: str,
content: str = None,
uniq_id: str = None,
tool_calls=None,
tool_call_id=None,
):
self.uniq_id = uniq_id if uniq_id is not None else str(uuid.uuid4()) self.uniq_id = uniq_id if uniq_id is not None else str(uuid.uuid4())
self.role = role self.role = role
self.content = content self.content = content
@@ -16,7 +23,7 @@ class Dialogue:
def __init__(self): def __init__(self):
self.dialogue: List[Message] = [] self.dialogue: List[Message] = []
# 获取当前时间 # 获取当前时间
self.current_time = datetime.now().strftime('%Y-%m-%d %H:%M:%S') self.current_time = datetime.now().strftime("%Y-%m-%d %H:%M:%S")
def put(self, message: Message): def put(self, message: Message):
self.dialogue.append(message) self.dialogue.append(message)
@@ -25,7 +32,15 @@ class Dialogue:
if m.tool_calls is not None: if m.tool_calls is not None:
dialogue.append({"role": m.role, "tool_calls": m.tool_calls}) dialogue.append({"role": m.role, "tool_calls": m.tool_calls})
elif m.role == "tool": elif m.role == "tool":
dialogue.append({"role": m.role, "tool_call_id": m.tool_call_id, "content": m.content}) dialogue.append(
{
"role": m.role,
"tool_call_id": (
str(uuid.uuid4()) if m.tool_call_id is None else m.tool_call_id
),
"content": m.content,
}
)
else: else:
dialogue.append({"role": m.role, "content": m.content}) dialogue.append({"role": m.role, "content": m.content})
@@ -44,23 +59,23 @@ class Dialogue:
else: else:
self.put(Message(role="system", content=new_content)) self.put(Message(role="system", content=new_content))
def get_llm_dialogue_with_memory(self, memory_str: str = None) -> List[Dict[str, str]]: def get_llm_dialogue_with_memory(
self, memory_str: str = None
) -> List[Dict[str, str]]:
if memory_str is None or len(memory_str) == 0: if memory_str is None or len(memory_str) == 0:
return self.get_llm_dialogue() return self.get_llm_dialogue()
# 构建带记忆的对话 # 构建带记忆的对话
dialogue = [] dialogue = []
# 添加系统提示和记忆 # 添加系统提示和记忆
system_message = next( system_message = next(
(msg for msg in self.dialogue if msg.role == "system"), None (msg for msg in self.dialogue if msg.role == "system"), None
) )
if system_message: if system_message:
enhanced_system_prompt = ( enhanced_system_prompt = (
f"{system_message.content}\n\n" f"{system_message.content}\n\n" f"相关记忆:\n{memory_str}"
f"相关记忆:\n{memory_str}"
) )
dialogue.append({"role": "system", "content": enhanced_system_prompt}) dialogue.append({"role": "system", "content": enhanced_system_prompt})
+567 -95
View File
@@ -2,35 +2,40 @@ import json
import socket import socket
import subprocess import subprocess
import re import re
import os
import numpy as np
import requests import requests
import opuslib_next
from pydub import AudioSegment
from typing import Dict, Any from typing import Dict, Any
from core.utils import tts, llm, intent, memory, vad, asr from core.utils import tts, llm, intent, memory, vad, asr
TAG = __name__ TAG = __name__
emoji_map = { emoji_map = {
'neutral': '😶', "neutral": "😶",
'happy': '🙂', "happy": "🙂",
'laughing': '😆', "laughing": "😆",
'funny': '😂', "funny": "😂",
'sad': '😔', "sad": "😔",
'angry': '😠', "angry": "😠",
'crying': '😭', "crying": "😭",
'loving': '😍', "loving": "😍",
'embarrassed': '😳', "embarrassed": "😳",
'surprised': '😲', "surprised": "😲",
'shocked': '😱', "shocked": "😱",
'thinking': '🤔', "thinking": "🤔",
'winking': '😉', "winking": "😉",
'cool': '😎', "cool": "😎",
'relaxed': '😌', "relaxed": "😌",
'delicious': '🤤', "delicious": "🤤",
'kissy': '😘', "kissy": "😘",
'confident': '😏', "confident": "😏",
'sleepy': '😴', "sleepy": "😴",
'silly': '😜', "silly": "😜",
'confused': '🙄' "confused": "🙄",
} }
def get_local_ip(): def get_local_ip():
try: try:
s = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) s = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
@@ -117,9 +122,9 @@ def is_punctuation_or_emoji(char):
"", # 中文顿号 "", # 中文顿号
"", "",
"", "",
"\"", # 中文双引号 + 英文引号 '"', # 中文双引号 + 英文引号
"", "",
":", # 中文冒号 + 英文冒号 ":", # 中文冒号 + 英文冒号
} }
if char.isspace() or char in punctuation_set: if char.isspace() or char in punctuation_set:
return True return True
@@ -345,20 +350,15 @@ def initialize_modules(
str(config.get("delete_audio", True)).lower() in ("true", "1", "yes"), str(config.get("delete_audio", True)).lower() in ("true", "1", "yes"),
) )
logger.bind(tag=TAG).info(f"初始化组件: asr成功 {select_asr_module}") logger.bind(tag=TAG).info(f"初始化组件: asr成功 {select_asr_module}")
# 初始化自定义prompt
if config.get("prompt", None) is not None:
modules["prompt"] = config["prompt"]
logger.bind(tag=TAG).info(f"初始化组件: prompt成功 {modules['prompt'][:50]}...")
return modules return modules
def analyze_emotion(text): def analyze_emotion(text):
""" """
分析文本情感并返回对应的emoji名称(支持中英文) 分析文本情感并返回对应的emoji名称(支持中英文)
""" """
if not text or not isinstance(text, str): if not text or not isinstance(text, str):
return 'neutral' return "neutral"
original_text = text original_text = text
text = text.lower().strip() text = text.lower().strip()
@@ -369,84 +369,444 @@ def analyze_emotion(text):
return emotion return emotion
# 标点符号分析 # 标点符号分析
has_exclamation = '!' in original_text or '' in original_text has_exclamation = "!" in original_text or "" in original_text
has_question = '?' in original_text or '' in original_text has_question = "?" in original_text or "" in original_text
has_ellipsis = '...' in original_text or '' in original_text has_ellipsis = "..." in original_text or "" in original_text
# 定义情感关键词映射(中英文扩展版) # 定义情感关键词映射(中英文扩展版)
emotion_keywords = { emotion_keywords = {
'happy': ['开心', '高兴', '快乐', '愉快', '幸福', '满意', '', '', '不错', '完美', '棒极了', '太好了', "happy": [
'好呀', '好的', 'happy', 'joy', 'great', 'good', 'nice', 'awesome', 'fantastic', 'wonderful'], "开心",
'laughing': ['哈哈', '哈哈哈', '呵呵', '嘿嘿', '嘻嘻', '笑死', '太好笑了', '笑死我了', 'lol', 'lmao', 'haha', "高兴",
'hahaha', 'hehe', 'rofl', 'funny', 'laugh'], "快乐",
'funny': ['搞笑', '滑稽', '', '幽默', '笑点', '段子', '笑话', '太逗了', 'hilarious', 'joke', 'comedy'], "愉快",
'sad': ['伤心', '难过', '悲哀', '悲伤', '忧郁', '郁闷', '沮丧', '失望', '想哭', '难受', '不开心', '', '呜呜', "幸福",
'sad', 'upset', 'unhappy', 'depressed', 'sorrow', 'gloomy'], "满意",
'angry': ['生气', '愤怒', '气死', '讨厌', '烦人', '可恶', '烦死了', '恼火', '暴躁', '火大', '愤怒', '气炸了', "",
'angry', 'mad', 'annoyed', 'furious', 'pissed', 'hate'], "",
'crying': ['哭泣', '泪流', '大哭', '伤心欲绝', '泪目', '流泪', '哭死', '哭晕', '想哭', '泪崩', "不错",
'cry', 'crying', 'tears', 'sob', 'weep'], "完美",
'loving': ['爱你', '喜欢', '', '亲爱的', '宝贝', '么么哒', '抱抱', '想你', '思念', '最爱', '亲亲', '喜欢你', "棒极了",
'love', 'like', 'adore', 'darling', 'sweetie', 'honey', 'miss you', 'heart'], "太好了",
'embarrassed': ['尴尬', '不好意思', '害羞', '脸红', '难为情', '社死', '丢脸', '出丑', "好呀",
'embarrassed', 'awkward', 'shy', 'blush'], "好的",
'surprised': ['惊讶', '吃惊', '天啊', '哇塞', '', '居然', '竟然', '没想到', '出乎意料', "happy",
'surprise', 'wow', 'omg', 'oh my god', 'amazing', 'unbelievable'], "joy",
'shocked': ['震惊', '吓到', '惊呆了', '不敢相信', '震撼', '吓死', '恐怖', '害怕', '吓人', "great",
'shocked', 'shocking', 'scared', 'frightened', 'terrified', 'horror'], "good",
'thinking': ['思考', '考虑', '想一下', '琢磨', '沉思', '冥想', '', '思考中', '在想', "nice",
'think', 'thinking', 'consider', 'ponder', 'meditate'], "awesome",
'winking': ['调皮', '眨眼', '你懂的', '坏笑', '邪恶', '奸笑', '使眼色', "fantastic",
'wink', 'teasing', 'naughty', 'mischievous'], "wonderful",
'cool': ['', '', '厉害', '棒极了', '真棒', '牛逼', '', '优秀', '杰出', '出色', '完美', ],
'cool', 'awesome', 'amazing', 'great', 'impressive', 'perfect'], "laughing": [
'relaxed': ['放松', '舒服', '惬意', '悠闲', '轻松', '舒适', '安逸', '自在', "哈哈",
'relax', 'relaxed', 'comfortable', 'cozy', 'chill', 'peaceful'], "哈哈哈",
'delicious': ['好吃', '美味', '', '', '可口', '香甜', '大餐', '大快朵颐', '流口水', '垂涎', "呵呵",
'delicious', 'yummy', 'tasty', 'yum', 'appetizing', 'mouthwatering'], "嘿嘿",
'kissy': ['亲亲', '么么', '', 'mua', 'muah', '亲一下', '飞吻', "嘻嘻",
'kiss', 'xoxo', 'hug', 'muah', 'smooch'], "笑死",
'confident': ['自信', '肯定', '确定', '毫无疑问', '当然', '必须的', '毫无疑问', '确信', '坚信', "太好笑了",
'confident', 'sure', 'certain', 'definitely', 'positive'], "笑死我了",
'sleepy': ['', '睡觉', '晚安', '想睡', '好累', '疲惫', '疲倦', '困了', '想休息', '睡意', "lol",
'sleep', 'sleepy', 'tired', 'exhausted', 'bedtime', 'good night'], "lmao",
'silly': ['', '', '', '', '', '', '憨憨', '傻乎乎', '呆萌', "haha",
'silly', 'stupid', 'dumb', 'foolish', 'goofy', 'ridiculous'], "hahaha",
'confused': ['疑惑', '不明白', '不懂', '困惑', '疑问', '为什么', '怎么回事', '啥意思', '不清楚', "hehe",
'confused', 'puzzled', 'doubt', 'question', 'what', 'why', 'how'] "rofl",
"funny",
"laugh",
],
"funny": [
"搞笑",
"滑稽",
"",
"幽默",
"笑点",
"段子",
"笑话",
"太逗了",
"hilarious",
"joke",
"comedy",
],
"sad": [
"伤心",
"难过",
"悲哀",
"悲伤",
"忧郁",
"郁闷",
"沮丧",
"失望",
"想哭",
"难受",
"不开心",
"",
"呜呜",
"sad",
"upset",
"unhappy",
"depressed",
"sorrow",
"gloomy",
],
"angry": [
"生气",
"愤怒",
"气死",
"讨厌",
"烦人",
"可恶",
"烦死了",
"恼火",
"暴躁",
"火大",
"愤怒",
"气炸了",
"angry",
"mad",
"annoyed",
"furious",
"pissed",
"hate",
],
"crying": [
"哭泣",
"泪流",
"大哭",
"伤心欲绝",
"泪目",
"流泪",
"哭死",
"哭晕",
"想哭",
"泪崩",
"cry",
"crying",
"tears",
"sob",
"weep",
],
"loving": [
"爱你",
"喜欢",
"",
"亲爱的",
"宝贝",
"么么哒",
"抱抱",
"想你",
"思念",
"最爱",
"亲亲",
"喜欢你",
"love",
"like",
"adore",
"darling",
"sweetie",
"honey",
"miss you",
"heart",
],
"embarrassed": [
"尴尬",
"不好意思",
"害羞",
"脸红",
"难为情",
"社死",
"丢脸",
"出丑",
"embarrassed",
"awkward",
"shy",
"blush",
],
"surprised": [
"惊讶",
"吃惊",
"天啊",
"哇塞",
"",
"居然",
"竟然",
"没想到",
"出乎意料",
"surprise",
"wow",
"omg",
"oh my god",
"amazing",
"unbelievable",
],
"shocked": [
"震惊",
"吓到",
"惊呆了",
"不敢相信",
"震撼",
"吓死",
"恐怖",
"害怕",
"吓人",
"shocked",
"shocking",
"scared",
"frightened",
"terrified",
"horror",
],
"thinking": [
"思考",
"考虑",
"想一下",
"琢磨",
"沉思",
"冥想",
"",
"思考中",
"在想",
"think",
"thinking",
"consider",
"ponder",
"meditate",
],
"winking": [
"调皮",
"眨眼",
"你懂的",
"坏笑",
"邪恶",
"奸笑",
"使眼色",
"wink",
"teasing",
"naughty",
"mischievous",
],
"cool": [
"",
"",
"厉害",
"棒极了",
"真棒",
"牛逼",
"",
"优秀",
"杰出",
"出色",
"完美",
"cool",
"awesome",
"amazing",
"great",
"impressive",
"perfect",
],
"relaxed": [
"放松",
"舒服",
"惬意",
"悠闲",
"轻松",
"舒适",
"安逸",
"自在",
"relax",
"relaxed",
"comfortable",
"cozy",
"chill",
"peaceful",
],
"delicious": [
"好吃",
"美味",
"",
"",
"可口",
"香甜",
"大餐",
"大快朵颐",
"流口水",
"垂涎",
"delicious",
"yummy",
"tasty",
"yum",
"appetizing",
"mouthwatering",
],
"kissy": [
"亲亲",
"么么",
"",
"mua",
"muah",
"亲一下",
"飞吻",
"kiss",
"xoxo",
"hug",
"muah",
"smooch",
],
"confident": [
"自信",
"肯定",
"确定",
"毫无疑问",
"当然",
"必须的",
"毫无疑问",
"确信",
"坚信",
"confident",
"sure",
"certain",
"definitely",
"positive",
],
"sleepy": [
"",
"睡觉",
"晚安",
"想睡",
"好累",
"疲惫",
"疲倦",
"困了",
"想休息",
"睡意",
"sleep",
"sleepy",
"tired",
"exhausted",
"bedtime",
"good night",
],
"silly": [
"",
"",
"",
"",
"",
"",
"憨憨",
"傻乎乎",
"呆萌",
"silly",
"stupid",
"dumb",
"foolish",
"goofy",
"ridiculous",
],
"confused": [
"疑惑",
"不明白",
"不懂",
"困惑",
"疑问",
"为什么",
"怎么回事",
"啥意思",
"不清楚",
"confused",
"puzzled",
"doubt",
"question",
"what",
"why",
"how",
],
} }
# 特殊句型判断(中英文) # 特殊句型判断(中英文)
# 赞美他人 # 赞美他人
if any(phrase in text for phrase in if any(
['你真', '你好', '您真', '你真棒', '你好厉害', '你太强了', '你真好', '你真聪明', phrase in text
'you are', 'you\'re', 'you look', 'you seem', 'so smart', 'so kind']): for phrase in [
return 'loving' "你真",
"你好",
"您真",
"你真棒",
"你好厉害",
"你太强了",
"你真好",
"你真聪明",
"you are",
"you're",
"you look",
"you seem",
"so smart",
"so kind",
]
):
return "loving"
# 自我赞美 # 自我赞美
if any(phrase in text for phrase in ['我真', '我最', '我太棒了', '我厉害', '我聪明', '我优秀', if any(
'i am', 'i\'m', 'i feel', 'so good', 'so happy']): phrase in text
return 'cool' for phrase in [
"我真",
"我最",
"我太棒了",
"我厉害",
"我聪明",
"我优秀",
"i am",
"i'm",
"i feel",
"so good",
"so happy",
]
):
return "cool"
# 晚安/睡觉相关 # 晚安/睡觉相关
if any(phrase in text for phrase in ['睡觉', '晚安', '睡了', '好梦', '休息了', '去睡了', if any(
'sleep', 'good night', 'bedtime', 'go to bed']): phrase in text
return 'sleepy' for phrase in [
"睡觉",
"晚安",
"睡了",
"好梦",
"休息了",
"去睡了",
"sleep",
"good night",
"bedtime",
"go to bed",
]
):
return "sleepy"
# 疑问句 # 疑问句
if has_question and not has_exclamation: if has_question and not has_exclamation:
return 'thinking' return "thinking"
# 强烈情感(感叹号) # 强烈情感(感叹号)
if has_exclamation and not has_question: if has_exclamation and not has_question:
# 检查是否是积极内容 # 检查是否是积极内容
positive_words = emotion_keywords['happy'] + emotion_keywords['laughing'] + emotion_keywords['cool'] positive_words = (
emotion_keywords["happy"]
+ emotion_keywords["laughing"]
+ emotion_keywords["cool"]
)
if any(word in text for word in positive_words): if any(word in text for word in positive_words):
return 'laughing' return "laughing"
# 检查是否是消极内容 # 检查是否是消极内容
negative_words = emotion_keywords['angry'] + emotion_keywords['sad'] + emotion_keywords['crying'] negative_words = (
emotion_keywords["angry"]
+ emotion_keywords["sad"]
+ emotion_keywords["crying"]
)
if any(word in text for word in negative_words): if any(word in text for word in negative_words):
return 'angry' return "angry"
return 'surprised' return "surprised"
# 省略号(表示犹豫或思考) # 省略号(表示犹豫或思考)
if has_ellipsis: if has_ellipsis:
return 'thinking' return "thinking"
# 关键词匹配(带权重) # 关键词匹配(带权重)
emotion_scores = {emotion: 0 for emotion in emoji_map.keys()} emotion_scores = {emotion: 0 for emotion in emoji_map.keys()}
@@ -466,18 +826,33 @@ def analyze_emotion(text):
# 根据分数选择最可能的情感 # 根据分数选择最可能的情感
max_score = max(emotion_scores.values()) max_score = max(emotion_scores.values())
if max_score == 0: if max_score == 0:
return 'happy' # 默认 return "happy" # 默认
# 可能有多个情感同分,根据上下文选择最合适的 # 可能有多个情感同分,根据上下文选择最合适的
top_emotions = [e for e, s in emotion_scores.items() if s == max_score] top_emotions = [e for e, s in emotion_scores.items() if s == max_score]
# 如果多个情感同分,使用以下优先级 # 如果多个情感同分,使用以下优先级
priority_order = [ priority_order = [
'laughing', 'crying', 'angry', 'surprised', 'shocked', # 强烈情感优先 "laughing",
'loving', 'happy', 'funny', 'cool', # 积极情感 "crying",
'sad', 'embarrassed', 'confused', # 消极情感 "angry",
'thinking', 'winking', 'relaxed', # 中性情感 "surprised",
'delicious', 'kissy', 'confident', 'sleepy', 'silly' # 特殊场景 "shocked", # 强烈情感优先
"loving",
"happy",
"funny",
"cool", # 积极情感
"sad",
"embarrassed",
"confused", # 消极情感
"thinking",
"winking",
"relaxed", # 中性情感
"delicious",
"kissy",
"confident",
"sleepy",
"silly", # 特殊场景
] ]
for emotion in priority_order: for emotion in priority_order:
@@ -485,3 +860,100 @@ def analyze_emotion(text):
return emotion return emotion
return top_emotions[0] # 如果都不在优先级列表里,返回第一个 return top_emotions[0] # 如果都不在优先级列表里,返回第一个
def audio_to_data(audio_file_path, is_opus=True):
# 获取文件后缀名
file_type = os.path.splitext(audio_file_path)[1]
if file_type:
file_type = file_type.lstrip(".")
# 读取音频文件,-nostdin 参数:不要从标准输入读取数据,否则FFmpeg会阻塞
audio = AudioSegment.from_file(
audio_file_path, format=file_type, parameters=["-nostdin"]
)
# 转换为单声道/16kHz采样率/16位小端编码(确保与编码器匹配)
audio = audio.set_channels(1).set_frame_rate(16000).set_sample_width(2)
# 音频时长(秒)
duration = len(audio) / 1000.0
# 获取原始PCM数据(16位小端)
raw_data = audio.raw_data
# 初始化Opus编码器
encoder = opuslib_next.Encoder(16000, 1, opuslib_next.APPLICATION_AUDIO)
# 编码参数
frame_duration = 60 # 60ms per frame
frame_size = int(16000 * frame_duration / 1000) # 960 samples/frame
datas = []
# 按帧处理所有音频数据(包括最后一帧可能补零)
for i in range(0, len(raw_data), frame_size * 2): # 16bit=2bytes/sample
# 获取当前帧的二进制数据
chunk = raw_data[i : i + frame_size * 2]
# 如果最后一帧不足,补零
if len(chunk) < frame_size * 2:
chunk += b"\x00" * (frame_size * 2 - len(chunk))
if is_opus:
# 转换为numpy数组处理
np_frame = np.frombuffer(chunk, dtype=np.int16)
# 编码Opus数据
frame_data = encoder.encode(np_frame.tobytes(), frame_size)
else:
frame_data = chunk if isinstance(chunk, bytes) else bytes(chunk)
datas.append(frame_data)
return datas, duration
def check_vad_update(before_config, new_config):
if (
new_config.get("selected_module") is None
or new_config["selected_module"].get("VAD") is None
):
return False
update_vad = False
current_vad_module = before_config["selected_module"]["VAD"]
new_vad_module = new_config["selected_module"]["VAD"]
current_vad_type = (
current_vad_module
if "type" not in before_config["VAD"][current_vad_module]
else before_config["VAD"][current_vad_module]["type"]
)
new_vad_type = (
new_vad_module
if "type" not in new_config["VAD"][new_vad_module]
else new_config["VAD"][new_vad_module]["type"]
)
print(f"前vad:{current_vad_type},后vad:{new_vad_type}")
update_vad = current_vad_type != new_vad_type
return update_vad
def check_asr_update(before_config, new_config):
if (
new_config.get("selected_module") is None
or new_config["selected_module"].get("ASR") is None
):
return False
update_asr = False
current_asr_module = before_config["selected_module"]["ASR"]
new_asr_module = new_config["selected_module"]["ASR"]
current_asr_type = (
current_asr_module
if "type" not in before_config["ASR"][current_asr_module]
else before_config["ASR"][current_asr_module]["type"]
)
new_asr_type = (
new_asr_module
if "type" not in new_config["ASR"][new_asr_module]
else new_config["ASR"][new_asr_module]["type"]
)
print(f"前asr:{current_asr_type},后asr:{new_asr_type}")
update_asr = current_asr_type != new_asr_type
return update_asr
+68 -9
View File
@@ -2,7 +2,8 @@ import asyncio
import websockets import websockets
from config.logger import setup_logging from config.logger import setup_logging
from core.connection import ConnectionHandler from core.connection import ConnectionHandler
from core.utils.util import initialize_modules from core.utils.util import initialize_modules, check_vad_update, check_asr_update
from config.config_loader import get_config_from_api
TAG = __name__ TAG = __name__
@@ -13,14 +14,21 @@ class WebSocketServer:
self.logger = setup_logging() self.logger = setup_logging()
self.config_lock = asyncio.Lock() self.config_lock = asyncio.Lock()
modules = initialize_modules( modules = initialize_modules(
self.logger, self.config, True, True, True, True, True, True self.logger,
self.config,
"VAD" in self.config["selected_module"],
"ASR" in self.config["selected_module"],
"LLM" in self.config["selected_module"],
"TTS" in self.config["selected_module"],
"Memory" in self.config["selected_module"],
"Intent" in self.config["selected_module"],
) )
self._vad = modules["vad"] self._vad = modules["vad"] if "vad" in modules else None
self._asr = modules["asr"] self._asr = modules["asr"] if "asr" in modules else None
self._tts = modules["tts"] self._tts = modules["tts"] if "tts" in modules else None
self._llm = modules["llm"] self._llm = modules["llm"] if "llm" in modules else None
self._intent = modules["intent"] self._intent = modules["intent"] if "intent" in modules else None
self._memory = modules["memory"] self._memory = modules["memory"] if "memory" in modules else None
self.active_connections = set() self.active_connections = set()
async def start(self): async def start(self):
@@ -44,7 +52,7 @@ class WebSocketServer:
self._tts, self._tts,
self._memory, self._memory,
self._intent, self._intent,
self # 传入当前 WebSocketServer 实例 self, # 传入server实例
) )
self.active_connections.add(handler) self.active_connections.add(handler)
try: try:
@@ -60,3 +68,54 @@ class WebSocketServer:
else: else:
# 如果是普通 HTTP 请求,返回 "server is running" # 如果是普通 HTTP 请求,返回 "server is running"
return websocket.respond(200, "Server is running\n") return websocket.respond(200, "Server is running\n")
async def update_config(self) -> bool:
"""更新服务器配置并重新初始化组件
Returns:
bool: 更新是否成功
"""
try:
async with self.config_lock:
# 重新获取配置
new_config = get_config_from_api(self.config)
if new_config is None:
self.logger.bind(tag=TAG).error("获取新配置失败")
return False
# 检查 VAD 和 ASR 类型是否需要更新
update_vad = check_vad_update(self.config, new_config)
update_asr = check_asr_update(self.config, new_config)
# 更新配置
self.config = new_config
# 重新初始化组件
modules = initialize_modules(
self.logger,
new_config,
update_vad,
update_asr,
"LLM" in new_config["selected_module"],
"TTS" in new_config["selected_module"],
"Memory" in new_config["selected_module"],
"Intent" in new_config["selected_module"],
)
# 更新组件实例
if "vad" in modules:
self._vad = modules["vad"]
if "asr" in modules:
self._asr = modules["asr"]
if "tts" in modules:
self._tts = modules["tts"]
if "llm" in modules:
self._llm = modules["llm"]
if "intent" in modules:
self._intent = modules["intent"]
if "memory" in modules:
self._memory = modules["memory"]
return True
except Exception as e:
self.logger.bind(tag=TAG).error(f"更新服务器配置失败: {str(e)}")
return False
+2 -1
View File
@@ -3,7 +3,8 @@
"在data目录下创建.mcp_server_settings.json文件,可以选择下面的MCP服务,也可以自行添加新的MCP服务。", "在data目录下创建.mcp_server_settings.json文件,可以选择下面的MCP服务,也可以自行添加新的MCP服务。",
"后面不断测试补充好用的mcp服务,欢迎大家一起补充。", "后面不断测试补充好用的mcp服务,欢迎大家一起补充。",
"记得删除注释行,des属性仅为说明,不会被解析。", "记得删除注释行,des属性仅为说明,不会被解析。",
"des和link属性,仅为说明安装方式,方便大家查看原始链接,不是必须项。" "des和link属性,仅为说明安装方式,方便大家查看原始链接,不是必须项。",
"当前支持stdio/sse两种模式。"
], ],
"mcpServers": { "mcpServers": {
"filesystem": { "filesystem": {

Some files were not shown because too many files have changed in this diff Show More