Compare commits

..
106 Commits
Author SHA1 Message Date
欣南科技andGitHub d59a769084 Merge pull request #1728 from xinnan-tech/redis-password
update:允许在docker-compose文件里设置redis密码
2025-07-01 16:35:20 +08:00
hrz 7a8fe1074c update:允许在docker-compose文件里设置redis密码 2025-07-01 16:34:37 +08:00
欣南科技andGitHub 3fc256ed9a Merge pull request #1727 from xinnan-tech/mcp-update
update:支持中文名称的mcp工具
2025-07-01 16:16:03 +08:00
hrz 18eed4417c update:支持中文名称的mcp工具 2025-07-01 16:07:35 +08:00
欣南科技andGitHub 20adbadc02 Merge pull request #1718 from xinnan-tech/newsnow
可配置新闻来源
2025-06-30 17:13:13 +08:00
hrz c68384be69 update:区分本地化部署和非本地化部署tts的配置字段 2025-06-30 16:56:56 +08:00
hrz 7619979046 update:简化新闻源配置 2025-06-30 14:24:46 +08:00
hrzandGitHub 54ccae8665 Merge pull request #1518 from shane04111/add_local_tts_change_voice
update: 添加本地tts支援面板的"角色音色"
2025-06-30 10:48:15 +08:00
hrzandGitHub 3015b01740 Merge branch 'newsnow' into add_local_tts_change_voice 2025-06-30 10:48:06 +08:00
hrzandGitHub 43d47a4f09 Merge pull request #1552 from liammazy/main
增加timeout配置项,单位为秒
2025-06-30 10:46:50 +08:00
hrzandGitHub 96abc34a92 Merge branch 'newsnow' into main 2025-06-30 10:44:39 +08:00
hrzandGitHub 26553b8875 Update openai.py 2025-06-30 10:41:15 +08:00
hrzandGitHub 5e3187e36b Merge pull request #1673 from goxofy/main
update: newsnow 插件支持前端配置新闻源  && fix issues/1658
2025-06-30 10:10:28 +08:00
hrzandGitHub eba8c3122f Merge branch 'newsnow' into main 2025-06-30 10:10:19 +08:00
hrzandGitHub f1a3f39782 Merge pull request #1712 from ifhaveif/main
🐞 fix: 只清空当前数据库,不清空整个数据库
2025-06-30 10:04:01 +08:00
xiongyonghui 0880f78002 🐞 fix: 只清空当前数据库,不清空整个数据库
-- emptyAll.lua
1、当redis被多个项目使用时,清空了其它项目的库(影响其它项目)
2025-06-30 00:17:39 +08:00
欣南科技andGitHub 1c16d497f9 Merge pull request #1706 from xinnan-tech/update-doc
update:增加MCP接入点的视频demo
2025-06-28 20:36:44 +08:00
hrz faaae95197 update:增加MCP接入点的视频demo 2025-06-28 20:36:05 +08:00
欣南科技andGitHub d5e0e8ab40 Merge pull request #1704 from xinnan-tech/py_websoket_link
update:优化超时资源优化
2025-06-27 23:50:25 +08:00
hrz 500b0ebde3 update:去除无关日志 2025-06-27 23:47:17 +08:00
hrz 6f6c39da23 update:优化超时资源优化 2025-06-27 23:38:05 +08:00
CGDandGitHub f8ccc7a92f Merge pull request #1703 from xinnan-tech/py_websoket_link
fix: 退出过快问题
2025-06-27 17:26:38 +08:00
Sakura-RanChen 95a678bb18 fix: 退出过快问题 2025-06-27 16:39:53 +08:00
欣南科技andGitHub e43ce6ba3c Merge pull request #1700 from xinnan-tech/manager-api-mcp
update:优化接入点连接跳转方式
2025-06-27 15:31:09 +08:00
hrz 7663b0b969 update:优化接入点连接跳转方式 2025-06-27 15:30:18 +08:00
欣南科技andGitHub ba898b0874 Merge pull request #1699 from xinnan-tech/manager-api-mcp
update:优化mcp地址的校验
2025-06-27 15:16:58 +08:00
hrz 04843010bc update:优化mcp地址的校验 2025-06-27 15:14:41 +08:00
欣南科技andGitHub 41c06efeee Merge pull request #1698 from xinnan-tech/manager-api-mcp
Manager api mcp
2025-06-27 14:59:27 +08:00
hrz d0425fa31a update:优化设备端读取mcp接入点工具 2025-06-27 14:57:49 +08:00
hrz 08753b97df update:页面显示接入点功能 2025-06-27 11:33:12 +08:00
hrz a86fef5f28 update:智能体获取mcp接入点工具列表 2025-06-27 10:06:40 +08:00
hrzandGitHub c5c2f59b31 Merge pull request #1688 from xinnan-tech/opt_tools
fix: 处理config_functions的类型转换,避免初始化错误
2025-06-26 18:32:00 +08:00
hrz 69ee8ab438 添加mcp工具列表接口 2025-06-26 18:29:07 +08:00
CGD 4a24bc8c21 fix: 处理config_functions的类型转换,避免初始化错误 2025-06-26 17:37:51 +08:00
hrz 1d293244d3 update:智控台增加mcp接入点的配置 2025-06-26 16:54:26 +08:00
hrz 6812c5ac40 update:更改版本号 2025-06-26 16:34:30 +08:00
TinKandGitHub eeb5ed920a Merge pull request #10 from goxofy/manual-add-device
fix
2025-06-26 16:29:29 +08:00
hrzandGitHub 90ad3e019f Merge pull request #1680 from xinnan-tech/opt_tools
MCP接入点完成
2025-06-26 16:27:42 +08:00
Tink 3af453c2df fix 2025-06-26 16:26:45 +08:00
hrz 3071d2bdfa update:完成mcp接入点对接 2025-06-26 16:25:39 +08:00
TinKandGitHub 90aad4f131 Merge pull request #9 from goxofy/manual-add-device
fix
2025-06-26 16:24:55 +08:00
Tink f69b1b34f1 fix 2025-06-26 16:22:21 +08:00
TinKandGitHub ece2c4b49f Merge pull request #8 from goxofy/manual-add-device
fix
2025-06-26 16:07:57 +08:00
Tink 033c8f9bee fix 2025-06-26 16:06:36 +08:00
TinKandGitHub 1b9ca95a0e Merge pull request #7 from goxofy/manual-add-device
UpdateWrapper
2025-06-26 16:00:31 +08:00
Tink 7e86a9f8a9 UpdateWrapper 2025-06-26 15:54:26 +08:00
TinKandGitHub 9dee744498 Merge pull request #6 from goxofy/manual-add-device
fix last_connected_at
2025-06-26 15:35:51 +08:00
Tink 32fa71e8f1 fix last_connected_at 2025-06-26 15:24:57 +08:00
hrz ec4694f859 update:优化服务端mcp 2025-06-26 14:35:12 +08:00
TinKandGitHub 717386f10a Merge pull request #5 from goxofy/manual-add-device
fix id in ai_device
2025-06-26 14:14:07 +08:00
Tink fe071fe41a fix id in ai_device 2025-06-26 14:13:29 +08:00
TinKandGitHub 560a7d083b Merge pull request #4 from goxofy/manual-add-device
manualAddDevice
2025-06-26 13:41:46 +08:00
Tink 4428c1a298 manualAddDevice 2025-06-26 13:40:01 +08:00
TinKandGitHub 760fd2fdc5 Merge pull request #3 from goxofy/manual-add-device
手动添加设备增加后端接口
2025-06-26 13:37:06 +08:00
Tink 84c11a281a 手动添加设备增加后端接口 2025-06-26 13:36:11 +08:00
TinKandGitHub ad682ca0c4 Merge pull request #2 from goxofy/manual-add-device
fix RequestService
2025-06-26 12:27:46 +08:00
Tink 2a349ca9ef fix RequestService 2025-06-26 12:25:19 +08:00
hrz da8435cfea update:视觉结果优化 2025-06-26 11:42:13 +08:00
hrz 0f36daf1fd update:优化工具回复 2025-06-26 11:27:23 +08:00
TinKandGitHub 1f3c90e9e5 Merge pull request #1 from goxofy/manual-add-device
manual-add-device
2025-06-26 11:16:23 +08:00
Tink 6ec7a1fe09 manual-add-device 2025-06-26 11:09:17 +08:00
hrz 2f78acaf4d update:优化iot操作 2025-06-26 09:34:00 +08:00
hrz 9412a26bfc update:优化工具目录结构 2025-06-26 09:28:29 +08:00
hrz 2348d9ceb3 update:统一工具注册及调用 2025-06-25 18:27:08 +08:00
JianYu Zheng 61b4b6617b 优化:提取共用方法,实现部分getAgentMcpToolsList内容
--AgentMcpAccessPointServiceImpl.java
1.获取URI对象,统一2个方法获取URI对象的异常处理
2.实现getAgentMcpToolsList部分内容
2025-06-25 17:39:01 +08:00
JianYu Zheng a8620f21a0 优化:方法分割提取
--AgentMcpAccessPointServiceImpl.java 方法分割提取共用
1.获取密钥方法
2.获取智能体mcp接入点url部分
3.获取对智能体id加密的token
2025-06-25 17:17:26 +08:00
Tink 78dc266eee change newsnow plugin doc 2025-06-25 16:49:54 +08:00
JianYu Zheng 0631fe37ea 添加新接口:获取智能体的Mcp接入点地址(/agent/mcp/address/{audioId})
--AgentMcpAccessPointServiceImpl.java 实现获取智能体的mcp接入点地址定义
--AgentMcpAccessPointController.java 添加获取智能体的Mcp接入点地址接口
2025-06-25 16:37:04 +08:00
Tink 8828f8401d newsnow 插件支持前端配置新闻源 2025-06-25 16:32:49 +08:00
JianYu Zheng 21339aba42 添加智能体mcp接入点接口方法定义
--AgentMcpAccessPointService.java
1.获取智能体的mcp接入点地址定义
2.获取智能体的mcp接入点已有的工具列表定义
2025-06-25 16:16:39 +08:00
JianYu Zheng 660231978f 添加一个哈希加密的工具类
--HashEncryptionUtil.java
1.指定哈希算法进行加密方法
2.使用md5进行加密
2025-06-25 16:15:06 +08:00
JianYu Zheng dde203b158 添加了一个新的系统参数,mcp接入点参数
--Constant.java mcp接入点参数的key常量
--SysParamsServiceImpl.java 添加了mcp参数验证方法,修改是进行验证
2025-06-25 11:21:41 +08:00
JianYu Zheng 0ee69fa2b2 添加了http发送工具类
--HttpSendUtils.java
添加了发送get请求,获取返回的body转换成字符串 方法
发送post请求,参数为json格式。取返回的body转换成字符串 方法
2025-06-25 11:06:36 +08:00
hrzandGitHub 8c1a8f7d55 Merge pull request #1663 from xinnan-tech/py_iot_fix
update:iot设备多指令适配
2025-06-24 16:52:10 +08:00
欣南科技andGitHub ef03b665c9 Merge pull request #1666 from xinnan-tech/aes_utils
update:新增AES加密方法,和python端AES一致
2025-06-24 16:48:56 +08:00
hrz b87906234d update:新增AES加密方法,和python端AES一致 2025-06-24 16:26:14 +08:00
CGD 896b318c49 update:iot设备多指令适配 2025-06-24 14:55:34 +08:00
CGDandGitHub b3d441c385 Merge pull request #1661 from xinnan-tech/web_fuction_MCP
update: mcp接入点页面实现
2025-06-24 14:41:12 +08:00
Sakura-RanChen 96de670bfe update: mcp接入点页面实现 2025-06-24 14:39:18 +08:00
欣南科技andGitHub bc2fc35cc9 Merge pull request #1642 from xinnan-tech/py_tts_listen
update:更新版本号
2025-06-20 17:42:53 +08:00
hrz 993d5395b2 update:更新版本号 2025-06-20 17:42:14 +08:00
欣南科技andGitHub 67f0b828ea Merge pull request #1641 from xinnan-tech/py_tts_listen
update:修复部分doubaoasr出现400错误的问题
2025-06-20 17:20:51 +08:00
hrz 879c1267b6 update:修复部分doubaoasr出现400错误的问题 2025-06-20 17:20:08 +08:00
hrzandGitHub 46c53e36a6 Merge pull request #1640 from xinnan-tech/py_tts_listen
fix: 反复打断任务没有被清除,vad四帧语音识别
2025-06-20 17:17:19 +08:00
hrz 09f6605cfe update:更新版本号 2025-06-20 17:17:09 +08:00
Sakura-RanChen 12cee4027a fix: 反复打断任务没有被清除,vad四帧语音识别 2025-06-20 16:25:08 +08:00
欣南科技andGitHub 979fea0d60 Merge pull request #1627 from xinnan-tech/update-remark
更新智控台两款意图识别的说明
2025-06-19 16:56:38 +08:00
hrz 94553c54bd 更新智控台两款意图识别的说明 2025-06-19 16:53:39 +08:00
hrzandGitHub 248db31c8b Merge pull request #1618 from xinnan-tech/py_bug_fix
update:添加MCP重连机制
2025-06-18 22:40:25 +08:00
hrzandGitHub 24bfa1ca15 Merge pull request #1615 from xinnan-tech/py_fix_type
fix: 豆包流式decode错误
2025-06-18 17:46:47 +08:00
欣南科技andGitHub 1e97a8febc Merge pull request #1620 from xinnan-tech/hot-fix
update:修复mcp返回json的bug
2025-06-18 17:28:51 +08:00
hrz 2742f2e1ff update:修复mcp返回json的bug 2025-06-18 17:28:21 +08:00
Sakura-RanChen 0a5ae70a7c 二次错误提醒 2025-06-18 16:32:58 +08:00
CGD 22d53bd36e update:添加MCP重连机制 2025-06-18 16:29:28 +08:00
Sakura-RanChen ebf68929ce update: 错误信息读取 2025-06-18 14:03:18 +08:00
Sakura-RanChen ef3b373211 fix: 豆包流式decode错误 2025-06-18 11:06:11 +08:00
hrzandGitHub 5012a51e1d Merge pull request #1609 from xinnan-tech/py_fix_type
Py fix type
2025-06-18 09:39:38 +08:00
Sakura-RanChen fb1f476a3c fix: 智控台切换默认视觉模型报错 2025-06-18 09:24:52 +08:00
Sakura-RanChen 1a31c8cd1d update: 视觉模块直接识别输出,不经过二次LLM 2025-06-17 16:13:24 +08:00
Sakura-RanChen d8bf5cdedf fix: openAi schema不支持list 2025-06-16 16:26:49 +08:00
shane0411 0ca1a6286c Merge remote-tracking branch 'upstream/main' into add_local_tts_change_voice 2025-06-12 12:36:13 +08:00
shane0411 a4aea694d4 Merge remote-tracking branch 'upstream/main' into add_local_tts_change_voice 2025-06-11 14:29:35 +08:00
Liam Mazy 0f31bfc78f 增加timeout配置项,单位为秒 2025-06-11 14:05:54 +08:00
shane0411 052d8c8962 feat: 添加音频路径音频文本的提示 2025-06-09 19:36:57 +08:00
shane0411 2eb5539c9b feat: 新增音频路径音频文本欄位,提供給本地tts選取音色
https://github.com/xinnan-tech/xiaozhi-esp32-server/issues/1503#issuecomment-2954952028
2025-06-09 19:18:58 +08:00
shane0411 be6fba5963 feat: 添加本地tts支援面板的"角色音色" 2025-06-09 15:20:27 +08:00
94 changed files with 4089 additions and 1048 deletions
+8 -3
View File
@@ -144,6 +144,11 @@
</a>
</td>
<td>
<a href="https://www.bilibili.com/video/BV1ZQKUzYExM" target="_blank">
<picture>
<img alt="MCP接入点" src="docs/images/demo13.png" />
</picture>
</a>
</td>
<td>
</td>
@@ -170,8 +175,8 @@
#### 🚀 部署方式选择
| 部署方式 | 特点 | 适用场景 | 部署文档 | 配置要求 | 视频教程 |
|---------|------|---------|---------|---------|---------|
| **最简化安装** | 智能对话、IOT、MCP、视觉感知,数据存储在配置文件 | 低配置环境,无需数据库 | [①Docker版](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%B8%80docker%E5%8F%AA%E8%BF%90%E8%A1%8Cserver) / [②源码部署](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%BA%8C%E6%9C%AC%E5%9C%B0%E6%BA%90%E7%A0%81%E5%8F%AA%E8%BF%90%E8%A1%8Cserver)| 如果使用`FunASR`要2核4G,如果全API,要2核2G | - |
| **全模块安装** | 智能对话、IOT、MCP、视觉感知、OTA、智控台,数据存储在数据库 | 完整功能体验 |[①Docker版](./docs/Deployment_all.md#%E6%96%B9%E5%BC%8F%E4%B8%80docker%E8%BF%90%E8%A1%8C%E5%85%A8%E6%A8%A1%E5%9D%97) / [②源码部署](./docs/Deployment_all.md#%E6%96%B9%E5%BC%8F%E4%BA%8C%E6%9C%AC%E5%9C%B0%E6%BA%90%E7%A0%81%E8%BF%90%E8%A1%8C%E5%85%A8%E6%A8%A1%E5%9D%97) / [③源码部署自动更新教程](./docs/dev-ops-integration.md) | 如果使用`FunASR`要4核8G,如果全API,要2核4G| [本地源码启动视频教程](https://www.bilibili.com/video/BV1wBJhz4Ewe) |
| **最简化安装** | 智能对话、IOT、MCP、视觉感知 | 低配置环境,数据存储在配置文件,无需数据库 | [①Docker版](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%B8%80docker%E5%8F%AA%E8%BF%90%E8%A1%8Cserver) / [②源码部署](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%BA%8C%E6%9C%AC%E5%9C%B0%E6%BA%90%E7%A0%81%E5%8F%AA%E8%BF%90%E8%A1%8Cserver)| 如果使用`FunASR`要2核4G,如果全API,要2核2G | - |
| **全模块安装** | 智能对话、IOT、MCP接入点、视觉感知、OTA、智控台 | 完整功能体验,数据存储在数据库 |[①Docker版](./docs/Deployment_all.md#%E6%96%B9%E5%BC%8F%E4%B8%80docker%E8%BF%90%E8%A1%8C%E5%85%A8%E6%A8%A1%E5%9D%97) / [②源码部署](./docs/Deployment_all.md#%E6%96%B9%E5%BC%8F%E4%BA%8C%E6%9C%AC%E5%9C%B0%E6%BA%90%E7%A0%81%E8%BF%90%E8%A1%8C%E5%85%A8%E6%A8%A1%E5%9D%97) / [③源码部署自动更新教程](./docs/dev-ops-integration.md) | 如果使用`FunASR`要4核8G,如果全API,要2核4G| [本地源码启动视频教程](https://www.bilibili.com/video/BV1wBJhz4Ewe) |
> 💡 提示:以下是按最新代码部署后的测试平台,有需要可烧录测试,并发为6个,每天会清空数据
@@ -226,7 +231,7 @@ Websocket接口地址: wss://2662r3426b.vicp.fun/xiaozhi/v1/
| 视觉感知系统 | 支持多种VLLM(视觉大模型),实现多模态交互 |
| 意图识别系统 | 支持LLM意图识别、Function Call函数调用,提供插件化意图处理机制 |
| 记忆系统 | 支持本地短期记忆、mem0ai接口记忆,具备记忆总结功能 |
| IOT/MCP控制协议 | 支持设备注册管理、智能控制接口,同时支持IOT、MCP控制协议 |
| 工具调用 | 支持客户端IOT协议、客户MCP协议、服务端MCP协议、MCP接入点协议、自定义工具函数 |
| 管理后台 | 提供Web管理界面,支持用户管理、系统配置和设备管理 |
| 测试工具 | 提供性能测试工具、视觉模型测试工具和音频交互测试工具 |
| 部署支持 | 支持Docker部署和本地部署,提供完整的配置文件管理 |
+3 -3
View File
@@ -42,9 +42,9 @@ http {
proxy_set_header Referer $http_referer;
proxy_set_header Cookie $http_cookie;
proxy_connect_timeout 10;
proxy_send_timeout 10;
proxy_read_timeout 10;
proxy_connect_timeout 15;
proxy_send_timeout 15;
proxy_read_timeout 15;
proxy_set_header X-Real-IP $remote_addr;
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
+1
View File
@@ -6,6 +6,7 @@ java -jar /app/xiaozhi-esp32-api.jar \
--spring.datasource.druid.username=${SPRING_DATASOURCE_DRUID_USERNAME} \
--spring.datasource.druid.password=${SPRING_DATASOURCE_DRUID_PASSWORD} \
--spring.data.redis.host=${SPRING_DATA_REDIS_HOST} \
--spring.data.redis.password=${SPRING_DATA_REDIS_PASSWORD} \
--spring.data.redis.port=${SPRING_DATA_REDIS_PORT} &
# 启动Nginx(前台运行保持容器存活)
Binary file not shown.

After

Width:  |  Height:  |  Size: 359 KiB

+134
View File
@@ -0,0 +1,134 @@
# MCP 接入点部署使用指南
本教程包含2个部分
- 1、如何开启mcp接入点
- 2、如何为智能体接入一个简单的mcp功能,如计算器功能
部署的前提条件:
- 1、你已经部署了全模块,因为mcp接入点需要全模块中的智控台功能
- 2、你想在不修改xiaozhi-server项目的前提下,扩展小智的功能
# 如何开启mcp接入点
## 第一步,下载mcp接入点项目源码
浏览器打开[mcp接入点项目地址](https://github.com/xinnan-tech/mcp-endpoint-server)
打开完,找到页面中一个绿色的按钮,写着`Code`的按钮,点开它,然后你就看到`Download ZIP`的按钮。
点击它,下载本项目源码压缩包。下载到你电脑后,解压它,此时它的名字可能叫`mcp-endpoint-server-main`
你需要把它重命名成`mcp-endpoint-server`
## 第二步,启动程序
这个项目是一个很简单的项目,建议使用docker运行。不过如果你不想使用docker运行,你可以参考[这个页面](https://github.com/xinnan-tech/mcp-endpoint-server/blob/main/README_dev.md)使用源码运行。以下是docker运行的方法
```
# 进入本项目源码根目录
cd mcp-endpoint-server
# 清除缓存
docker compose -f docker-compose.yml down
docker stop mcp-endpoint-server
docker rm mcp-endpoint-server
docker rmi ghcr.nju.edu.cn/xinnan-tech/mcp-endpoint-server:latest
# 启动docker容器
docker compose -f docker-compose.yml up -d
# 查看日志
docker logs -f mcp-endpoint-server
```
此时,日志里会输出类似以下的日志
```
======================================================
接口地址: http://172.1.1.1:8004/mcp_endpoint/health?key=xxxx
=======上面的地址是MCP接入点地址,请勿泄露给任何人============
```
请你把接口地址复制出来:
由于你是docker部署,切不可直接使用上面的地址!
由于你是docker部署,切不可直接使用上面的地址!
由于你是docker部署,切不可直接使用上面的地址!
你先把地址复制出来,放在一个草稿里,你要知道你的电脑的局域网ip是什么,例如我的电脑局域网ip是`192.168.1.25`,那么
原来我的接口地址
```
http://172.1.1.1:8004/mcp_endpoint/health?key=xxxx
```
就要改成
```
http://192.168.1.25:8004/mcp_endpoint/health?key=xxxx
```
改好后,请使用浏览器直接访问这个接口。当浏览器出现类似这样的代码,说明是成功了。
```
{"result":{"status":"success","connections":{"tool_connections":0,"robot_connections":0,"total_connections":0}},"error":null,"id":null,"jsonrpc":"2.0"}
```
请你保留好这个`接口地址`,下一步要用到。
## 第三步,配置智控台
使用管理员账号,登录智控台,点击顶部`参数字典`,选择`参数管理`功能。
然后搜索参数`server.mcp_endpoint`,此时,它的值应该是`null`值。
点击修改按钮,把上一步得来的`接口地址`粘贴到`参数值`里。然后保存。
如果能保存成功,说明一切顺利,你可以去智能体查看效果了。如果不成功,说明智控台无法访问mcp接入点,很大概率是网络防火墙,或者没有填写正确的局域网ip。
# 如何为智能体接入一个简单的mcp功能,如计算器功能
如果以上步骤顺利,你可以进入智能体管理,点击`配置角色`,在`意图识别`的右边,有一个`编辑功能`的按钮。
点击这个按钮。在弹出的页面里,位于底部,会有`MCP接入点`,正常来说,会显示这个智能体的`MCP接入点地址`,接下来,我们来给这个智能体扩展一个基于MCP技术的计算器的功能。
这个`MCP接入点地址`很重要,你等一下会用到。
## 第一步 下载虾哥MCP计算器项目代码
浏览器打开虾哥写的[计算器项目](https://github.com/78/mcp-calculator)
打开完,找到页面中一个绿色的按钮,写着`Code`的按钮,点开它,然后你就看到`Download ZIP`的按钮。
点击它,下载本项目源码压缩包。下载到你电脑后,解压它,此时它的名字可能叫`mcp-calculatorr-main`
你需要把它重命名成`mcp-calculator`。接下来,我们用命令行进入项目目录即安装依赖
```bash
# 进入项目目录
cd mcp-calculator
conda remove -n mcp-calculator --all -y
conda create -n mcp-calculator python=3.10 -y
conda activate mcp-calculator
pip install -r requirements.txt
```
## 第二步 启动
启动前,先从你的智控台的智能体里,复制到了MCP接入点的地址。
例如我的智能体的mcp地址是
```
ws://192.168.4.7:8004/mcp_endpoint/mcp/?token=abc
```
开始输入命令
```bash
export MCP_ENDPOINT=ws://192.168.4.7:8004/mcp_endpoint/mcp/?token=abc
```
输入完后,启动程序
```bash
python mcp_pipe.py calculator.py
```
启动完后,你再进入智控台,点击刷新MCP的接入状态,就会看到你扩展的功能列表了。
+105
View File
@@ -0,0 +1,105 @@
# get_news_from_newsnow 插件新闻源配置指南
## 概述
`get_news_from_newsnow` 插件现在支持通过Web管理界面动态配置新闻源,不再需要修改代码。用户可以在智控台中为每个智能体配置不同的新闻源。
## 配置方式
### 1. 通过Web管理界面配置(推荐)
1. 登录智控台
2. 进入"角色配置"页面
3. 选择要配置的智能体
4. 点击"编辑功能"按钮
5. 在右侧参数配置区域找到"newsnow新闻聚合"插件
6. 在"新闻源配置"字段中输入分号分隔的中文名称
### 2. 配置文件方式
`config.yaml` 中配置:
```yaml
plugins:
get_news_from_newsnow:
url: "https://newsnow.busiyi.world/api/s?id="
news_sources: "澎湃新闻;百度热搜;财联社;微博;抖音"
```
## 新闻源配置格式
新闻源配置使用分号分隔的中文名称,格式为:
```
中文名称1;中文名称2;中文名称3
```
### 配置示例
```
澎湃新闻;百度热搜;财联社;微博;抖音;知乎;36氪
```
## 支持的新闻源
插件支持以下新闻源的中文名称:
- 澎湃新闻
- 百度热搜
- 财联社
- 微博
- 抖音
- 知乎
- 36氪
- 华尔街见闻
- IT之家
- 今日头条
- 虎扑
- 哔哩哔哩
- 快手
- 雪球
- 格隆汇
- 法布财经
- 金十数据
- 牛客
- 少数派
- 稀土掘金
- 凤凰网
- 虫部落
- 联合早报
- 酷安
- 远景论坛
- 参考消息
- 卫星通讯社
- 百度贴吧
- 靠谱新闻
- 以及更多...
## 默认配置
如果未配置新闻源,插件将使用以下默认配置:
```
澎湃新闻;百度热搜;财联社
```
## 使用说明
1. **配置新闻源**:在Web界面或配置文件中设置新闻源的中文名称,用分号分隔
2. **调用插件**:用户可以说"播报新闻"或"获取新闻"
3. **指定新闻源**:用户可以说"播报澎湃新闻"或"获取百度热搜"
4. **获取详情**:用户可以说"详细介绍这条新闻"
## 工作原理
1. 插件接受中文名称作为参数(如"澎湃新闻")
2. 根据配置的新闻源列表,将中文名称转换为对应的英文ID(如"thepaper"
3. 使用英文ID调用API获取新闻数据
4. 返回新闻内容给用户
## 注意事项
1. 配置的中文名称必须与 CHANNEL_MAP 中定义的名称完全一致
2. 配置更改后需要重启服务或重新加载配置
3. 如果配置的新闻源无效,插件会自动使用默认新闻源
4. 多个新闻源之间使用英文分号(;)分隔,不要使用中文分号(;)
@@ -111,6 +111,11 @@ public interface Constant {
*/
String FILE_EXTENSION_SEG = ".";
/**
* mcp接入点路径
*/
String SERVER_MCP_ENDPOINT = "server.mcp_endpoint";
/**
* 无记忆
*/
@@ -227,7 +232,7 @@ public interface Constant {
/**
* 版本号
*/
public static final String VERSION = "0.5.7";
public static final String VERSION = "0.6.1";
/**
* 无效固件URL
@@ -0,0 +1,78 @@
package xiaozhi.common.utils;
import java.nio.charset.StandardCharsets;
import java.util.Base64;
import javax.crypto.Cipher;
import javax.crypto.spec.SecretKeySpec;
public class AESUtils {
private static final String ALGORITHM = "AES";
private static final String TRANSFORMATION = "AES/ECB/PKCS5Padding";
/**
* AES加密
*
* @param key 密钥(16位、24位或32位)
* @param plainText 待加密字符串
* @return 加密后的Base64字符串
*/
public static String encrypt(String key, String plainText) {
try {
// 确保密钥长度为16、24或32位
byte[] keyBytes = padKey(key.getBytes(StandardCharsets.UTF_8));
SecretKeySpec secretKey = new SecretKeySpec(keyBytes, ALGORITHM);
Cipher cipher = Cipher.getInstance(TRANSFORMATION);
cipher.init(Cipher.ENCRYPT_MODE, secretKey);
byte[] encryptedBytes = cipher.doFinal(plainText.getBytes(StandardCharsets.UTF_8));
return Base64.getEncoder().encodeToString(encryptedBytes);
} catch (Exception e) {
throw new RuntimeException("AES加密失败", e);
}
}
/**
* AES解密
*
* @param key 密钥(16位、24位或32位)
* @param encryptedText 待解密的Base64字符串
* @return 解密后的字符串
*/
public static String decrypt(String key, String encryptedText) {
try {
// 确保密钥长度为16、24或32位
byte[] keyBytes = padKey(key.getBytes(StandardCharsets.UTF_8));
SecretKeySpec secretKey = new SecretKeySpec(keyBytes, ALGORITHM);
Cipher cipher = Cipher.getInstance(TRANSFORMATION);
cipher.init(Cipher.DECRYPT_MODE, secretKey);
byte[] encryptedBytes = Base64.getDecoder().decode(encryptedText);
byte[] decryptedBytes = cipher.doFinal(encryptedBytes);
return new String(decryptedBytes, StandardCharsets.UTF_8);
} catch (Exception e) {
throw new RuntimeException("AES解密失败", e);
}
}
/**
* 填充密钥到指定长度(16、24或32位)
*
* @param keyBytes 原始密钥字节数组
* @return 填充后的密钥字节数组
*/
private static byte[] padKey(byte[] keyBytes) {
int keyLength = keyBytes.length;
if (keyLength == 16 || keyLength == 24 || keyLength == 32) {
return keyBytes;
}
// 如果密钥长度不足,用0填充;如果超过,截取前32位
byte[] paddedKey = new byte[32];
System.arraycopy(keyBytes, 0, paddedKey, 0, Math.min(keyLength, 32));
return paddedKey;
}
}
@@ -0,0 +1,52 @@
package xiaozhi.common.utils;
import lombok.extern.slf4j.Slf4j;
import java.security.MessageDigest;
import java.security.NoSuchAlgorithmException;
/**
* 哈希加密算法的工具类
* @author zjy
*/
@Slf4j
public class HashEncryptionUtil {
/**
* 使用md5进行加密
* @param context 被加密的内容
* @return 哈希值
*/
public static String Md5hexDigest(String context){
return hexDigest(context,"MD5");
}
/**
* 指定哈希算法进行加密
* @param context 被加密的内容
* @param algorithm 哈希算法
* @return 哈希值
*/
public static String hexDigest(String context,String algorithm ){
// 获取MD5算法实例
MessageDigest md = null;
try {
md = MessageDigest.getInstance(algorithm);
} catch (NoSuchAlgorithmException e) {
log.error("加密失败的算法:{}",algorithm);
throw new RuntimeException("加密失败,"+ algorithm +"哈希算法系统不支持");
}
// 计算智能体id的MD5值
byte[] messageDigest = md.digest(context.getBytes());
// 将字节数组转换为十六进制字符串
StringBuilder hexString = new StringBuilder();
for (byte b : messageDigest) {
String hex = Integer.toHexString(0xFF & b);
if (hex.length() == 1) {
hexString.append('0');
}
hexString.append(hex);
}
return hexString.toString();
}
}
@@ -0,0 +1,66 @@
package xiaozhi.modules.agent.controller;
import java.util.List;
import org.apache.shiro.authz.annotation.RequiresPermissions;
import org.springframework.web.bind.annotation.GetMapping;
import org.springframework.web.bind.annotation.PathVariable;
import org.springframework.web.bind.annotation.RequestMapping;
import org.springframework.web.bind.annotation.RestController;
import io.swagger.v3.oas.annotations.Operation;
import io.swagger.v3.oas.annotations.tags.Tag;
import lombok.RequiredArgsConstructor;
import xiaozhi.common.user.UserDetail;
import xiaozhi.common.utils.Result;
import xiaozhi.modules.agent.service.AgentMcpAccessPointService;
import xiaozhi.modules.agent.service.AgentService;
import xiaozhi.modules.security.user.SecurityUser;
@Tag(name = "智能体Mcp接入点管理")
@RequiredArgsConstructor
@RestController
@RequestMapping("/agent/mcp")
public class AgentMcpAccessPointController {
private final AgentMcpAccessPointService agentMcpAccessPointService;
private final AgentService agentService;
/**
* 获取智能体的Mcp接入点地址
*
* @param audioId 智能体id
* @return 返回错误提醒或者Mcp接入点地址
*/
@Operation(summary = "获取智能体的Mcp接入点地址")
@GetMapping("/address/{agentId}")
@RequiresPermissions("sys:role:normal")
public Result<String> getAgentMcpAccessAddress(@PathVariable("agentId") String agentId) {
// 获取当前用户
UserDetail user = SecurityUser.getUser();
// 检查权限
if (!agentService.checkAgentPermission(agentId, user.getId())) {
return new Result<String>().error("没有权限查看该智能体的MCP接入点地址");
}
String agentMcpAccessAddress = agentMcpAccessPointService.getAgentMcpAccessAddress(agentId);
if (agentMcpAccessAddress == null) {
return new Result<String>().ok("请联系管理员进入参数管理配置mcp接入点地址");
}
return new Result<String>().ok(agentMcpAccessAddress);
}
@Operation(summary = "获取智能体的Mcp工具列表")
@GetMapping("/tools/{agentId}")
@RequiresPermissions("sys:role:normal")
public Result<List<String>> getAgentMcpToolsList(@PathVariable("agentId") String agentId) {
// 获取当前用户
UserDetail user = SecurityUser.getUser();
// 检查权限
if (!agentService.checkAgentPermission(agentId, user.getId())) {
return new Result<List<String>>().error("没有权限查看该智能体的MCP工具列表");
}
List<String> agentMcpToolsList = agentMcpAccessPointService.getAgentMcpToolsList(agentId);
return new Result<List<String>>().ok(agentMcpToolsList);
}
}
@@ -0,0 +1,32 @@
package xiaozhi.modules.agent.dto;
import lombok.Data;
/**
* MCP JSON-RPC 请求 DTO
*/
@Data
public class McpJsonRpcRequest {
private String jsonrpc = "2.0";
private String method;
private Object params;
private Integer id;
public McpJsonRpcRequest() {
}
public McpJsonRpcRequest(String method) {
this.method = method;
}
public McpJsonRpcRequest(String method, Object params, Integer id) {
this.method = method;
this.params = params;
this.id = id;
}
public McpJsonRpcRequest(String method, Object params) {
this.method = method;
this.params = params;
}
}
@@ -0,0 +1,48 @@
package xiaozhi.modules.agent.dto;
import lombok.Data;
/**
* MCP JSON-RPC 响应 DTO
*/
@Data
public class McpJsonRpcResponse {
private String jsonrpc = "2.0";
private Integer id;
private McpResult result;
private McpError error;
public McpJsonRpcResponse() {
}
@Data
public static class McpResult {
private String type;
private String message;
private String agent_id;
private McpTool[] tools;
public McpResult() {
}
}
@Data
public static class McpTool {
private String name;
private String description;
private Object inputSchema;
public McpTool() {
}
}
@Data
public static class McpError {
private Integer code;
private String message;
private Object data;
public McpError() {
}
}
}
@@ -0,0 +1,25 @@
package xiaozhi.modules.agent.service;
import java.util.List;
/**
* 智能体Mcp接入点处理service
*
* @author zjy
*/
public interface AgentMcpAccessPointService {
/**
* 获取智能体的mcp接入点地址
* @param id 智能体id
* @return mcp接入点地址
*/
String getAgentMcpAccessAddress(String id);
/**
* 获取智能体的mcp接入点已有的工具列表
* @param id 智能体id
* @return 工具列表
*/
List<String> getAgentMcpToolsList(String id);
}
@@ -19,6 +19,8 @@ import xiaozhi.modules.agent.service.AgentChatAudioService;
import xiaozhi.modules.agent.service.AgentChatHistoryService;
import xiaozhi.modules.agent.service.AgentService;
import xiaozhi.modules.agent.service.biz.AgentChatHistoryBizService;
import xiaozhi.modules.device.entity.DeviceEntity;
import xiaozhi.modules.device.service.DeviceService;
/**
* {@link AgentChatHistoryBizService} impl
@@ -35,6 +37,7 @@ public class AgentChatHistoryBizServiceImpl implements AgentChatHistoryBizServic
private final AgentChatHistoryService agentChatHistoryService;
private final AgentChatAudioService agentChatAudioService;
private final RedisUtils redisUtils;
private final DeviceService deviceService;
/**
* 处理聊天记录上报,包括文件上传和相关信息记录
@@ -68,6 +71,15 @@ public class AgentChatHistoryBizServiceImpl implements AgentChatHistoryBizServic
// 更新设备最后对话时间
redisUtils.set(RedisKeys.getAgentDeviceLastConnectedAtById(agentId), new Date());
// 更新设备最后连接时间
DeviceEntity device = deviceService.getDeviceByMacAddress(macAddress);
if (device != null) {
deviceService.updateDeviceConnectionInfo(agentId, device.getId(), null);
} else {
log.warn("聊天记录上报时,未找到mac地址为 {} 的设备", macAddress);
}
return Boolean.TRUE;
}
@@ -0,0 +1,207 @@
package xiaozhi.modules.agent.service.impl;
import java.net.URI;
import java.net.URISyntaxException;
import java.net.URLEncoder;
import java.nio.charset.StandardCharsets;
import java.util.List;
import java.util.Map;
import java.util.concurrent.TimeUnit;
import java.util.stream.Collectors;
import org.apache.commons.lang3.StringUtils;
import org.springframework.stereotype.Service;
import lombok.AllArgsConstructor;
import lombok.extern.slf4j.Slf4j;
import xiaozhi.common.constant.Constant;
import xiaozhi.common.utils.AESUtils;
import xiaozhi.common.utils.HashEncryptionUtil;
import xiaozhi.common.utils.JsonUtils;
import xiaozhi.modules.agent.dto.McpJsonRpcRequest;
import xiaozhi.modules.agent.service.AgentMcpAccessPointService;
import xiaozhi.modules.sys.service.SysParamsService;
import xiaozhi.modules.sys.utils.WebSocketClientManager;
@AllArgsConstructor
@Service
@Slf4j
public class AgentMcpAccessPointServiceImpl implements AgentMcpAccessPointService {
private SysParamsService sysParamsService;
@Override
public String getAgentMcpAccessAddress(String id) {
// 获取到mcp的地址
String url = sysParamsService.getValue(Constant.SERVER_MCP_ENDPOINT, true);
if (StringUtils.isBlank(url) || "null".equals(url)) {
return null;
}
URI uri = getURI(url);
// 获取智能体mcp的url前缀
String agentMcpUrl = getAgentMcpUrl(uri);
// 获取密钥
String key = getSecretKey(uri);
// 获取加密的token
String encryptToken = encryptToken(id, key);
// 对token进行URL编码
String encodedToken = URLEncoder.encode(encryptToken, StandardCharsets.UTF_8);
// 返回智能体Mcp路径的格式
agentMcpUrl = "%s/mcp/?token=%s".formatted(agentMcpUrl, encodedToken);
return agentMcpUrl;
}
@Override
public List<String> getAgentMcpToolsList(String id) {
String wsUrl = getAgentMcpAccessAddress(id);
if (StringUtils.isBlank(wsUrl)) {
return List.of();
}
// 将 /mcp 替换为 /call
wsUrl = wsUrl.replace("/mcp/", "/call/");
try {
// 创建 WebSocket 连接
try (WebSocketClientManager client = WebSocketClientManager.build(
new WebSocketClientManager.Builder()
.uri(wsUrl)
.connectTimeout(5, TimeUnit.SECONDS)
.maxSessionDuration(9, TimeUnit.SECONDS))) {
// 发送初始化消息
McpJsonRpcRequest initializeRequest = new McpJsonRpcRequest("initialize",
Map.of(
"protocolVersion", "2024-11-05",
"capabilities", Map.of(
"roots", Map.of("listChanged", false),
"sampling", Map.of()),
"clientInfo", Map.of(
"name", "xz-mcp-broker",
"version", "0.0.1")),
1);
client.sendJson(initializeRequest);
// 等待初始化响应
Thread.sleep(200);
// 发送初始化完成通知
// 对于通知类型的消息,手动构建JSON以避免包含null字段
String notificationJson = "{\"jsonrpc\":\"2.0\",\"method\":\"notifications/initialized\"}";
client.sendText(notificationJson);
// 等待 0.2 秒
Thread.sleep(200);
// 发送工具列表请求
McpJsonRpcRequest toolsRequest = new McpJsonRpcRequest("tools/list", null, 2);
client.sendJson(toolsRequest);
// 监听响应,直到收到包含 id=2 的响应(tools/list响应)
List<String> responses = client.listener(response -> {
try {
// 先尝试解析为通用JSON对象来获取id
Map<String, Object> jsonMap = JsonUtils.parseObject(response, Map.class);
return jsonMap != null && Integer.valueOf(2).equals(jsonMap.get("id"));
} catch (Exception e) {
log.warn("解析响应失败: {}", response, e);
return false;
}
});
// 处理响应
for (String response : responses) {
try {
// 先解析为通用JSON对象
Map<String, Object> jsonMap = JsonUtils.parseObject(response, Map.class);
if (jsonMap != null && Integer.valueOf(2).equals(jsonMap.get("id"))) {
// 检查是否有result字段
Object resultObj = jsonMap.get("result");
if (resultObj instanceof Map) {
Map<String, Object> resultMap = (Map<String, Object>) resultObj;
Object toolsObj = resultMap.get("tools");
if (toolsObj instanceof List) {
List<Map<String, Object>> toolsList = (List<Map<String, Object>>) toolsObj;
// 提取工具名称列表
return toolsList.stream()
.map(tool -> (String) tool.get("name"))
.filter(name -> name != null)
.collect(Collectors.toList());
}
}
}
} catch (Exception e) {
log.warn("处理工具列表响应失败: {}", response, e);
}
}
log.warn("未找到有效的工具列表响应");
return List.of();
}
} catch (Exception e) {
log.error("获取智能体 MCP 工具列表失败,智能体ID: {}", id, e);
return List.of();
}
}
/**
* 获取URI对象
*
* @param url 路径
* @return URI对象
*/
private static URI getURI(String url) {
try {
return new URI(url);
} catch (URISyntaxException e) {
log.error("路径格式不正确路径:{}\n错误信息:{}", url, e.getMessage());
throw new RuntimeException("mcp的地址存在错误,请进入参数管理修改mcp接入点地址");
}
}
/**
* 获取密钥
*
* @param uri mcp地址
* @return 密钥
*/
private static String getSecretKey(URI uri) {
// 获取参数
String query = uri.getQuery();
// 获取aes加密密钥
String str = "key=";
return query.substring(query.indexOf(str) + str.length());
}
/**
* 获取智能体mcp接入点url
*
* @param uri mcp地址
* @return 智能体mcp接入点url
*/
private String getAgentMcpUrl(URI uri) {
// 获取协议
String wsScheme = (uri.getScheme().equals("https")) ? "wss" : "ws";
// 获取主机,端口,路径
String path = uri.getSchemeSpecificPart();
// 获取到最后一个/前的path
path = path.substring(0, path.lastIndexOf("/"));
return wsScheme + ":" + path;
}
/**
* 获取对智能体id加密的token
*
* @param agentId 智能体id
* @param key 加密密钥
* @return 加密后token
*/
private static String encryptToken(String agentId, String key) {
// 使用md5对智能体id进行加密
String md5 = HashEncryptionUtil.Md5hexDigest(agentId);
// aes需要加密文本
String json = "{\"agentId\": \"%s\"}".formatted(md5);
// 加密后成token值
return AESUtils.encrypt(key, json);
}
}
@@ -56,6 +56,9 @@ public class AgentTemplateServiceImpl extends ServiceImpl<AgentTemplateDao, Agen
wrapper.set("tts_model_id", modelId);
wrapper.set("tts_voice_id", null);
break;
case "VLLM":
wrapper.set("vllm_model_id", modelId);
break;
case "MEMORY":
wrapper.set("mem_model_id", modelId);
break;
@@ -1,6 +1,10 @@
package xiaozhi.modules.config.service.impl;
import java.util.*;
import java.util.ArrayList;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
import java.util.Objects;
import org.apache.commons.lang3.StringUtils;
import org.springframework.stereotype.Service;
@@ -15,6 +19,7 @@ import xiaozhi.common.utils.JsonUtils;
import xiaozhi.modules.agent.entity.AgentEntity;
import xiaozhi.modules.agent.entity.AgentPluginMapping;
import xiaozhi.modules.agent.entity.AgentTemplateEntity;
import xiaozhi.modules.agent.service.AgentMcpAccessPointService;
import xiaozhi.modules.agent.service.AgentPluginMappingService;
import xiaozhi.modules.agent.service.AgentService;
import xiaozhi.modules.agent.service.AgentTemplateService;
@@ -39,6 +44,7 @@ public class ConfigServiceImpl implements ConfigService {
private final RedisUtils redisUtils;
private final TimbreService timbreService;
private final AgentPluginMappingService agentPluginMappingService;
private final AgentMcpAccessPointService agentMcpAccessPointService;
@Override
public Object getConfig(Boolean isCache) {
@@ -66,6 +72,8 @@ public class ConfigServiceImpl implements ConfigService {
null,
null,
null,
null,
null,
agent.getVadModelId(),
agent.getAsrModelId(),
null,
@@ -102,9 +110,13 @@ public class ConfigServiceImpl implements ConfigService {
}
// 获取音色信息
String voice = null;
String referenceAudio = null;
String referenceText = null;
TimbreDetailsVO timbre = timbreService.get(agent.getTtsVoiceId());
if (timbre != null) {
voice = timbre.getTtsVoice();
referenceAudio = timbre.getReferenceAudio();
referenceText = timbre.getReferenceText();
}
// 构建返回数据
Map<String, Object> result = new HashMap<>();
@@ -144,6 +156,12 @@ public class ConfigServiceImpl implements ConfigService {
result.put("plugins", pluginParams);
}
}
// 获取mcp接入点地址
String mcpEndpoint = agentMcpAccessPointService.getAgentMcpAccessAddress(agent.getId());
if (StringUtils.isNotBlank(mcpEndpoint) && mcpEndpoint.startsWith("ws")) {
mcpEndpoint = mcpEndpoint.replace("/mcp/", "/call/");
result.put("mcp_endpoint", mcpEndpoint);
}
// 构建模块配置
buildModuleConfig(
@@ -151,6 +169,8 @@ public class ConfigServiceImpl implements ConfigService {
agent.getSystemPrompt(),
agent.getSummaryMemory(),
voice,
referenceAudio,
referenceText,
agent.getVadModelId(),
agent.getAsrModelId(),
agent.getLlmModelId(),
@@ -167,7 +187,7 @@ public class ConfigServiceImpl implements ConfigService {
/**
* 构建配置信息
*
* @param paramsList 系统参数列表
* @param config 系统参数列表
* @return 配置信息
*/
private Object buildConfig(Map<String, Object> config) {
@@ -238,21 +258,25 @@ public class ConfigServiceImpl implements ConfigService {
/**
* 构建模块配置
*
* @param prompt 提示词
* @param voice 音色
* @param vadModelId VAD模型ID
* @param asrModelId ASR模型ID
* @param llmModelId LLM模型ID
* @param ttsModelId TTS模型ID
* @param memModelId 记忆模型ID
* @param intentModelId 意图模型ID
* @param result 结果Map
* @param prompt 提示词
* @param voice 音色
* @param referenceAudio 参考音频路径
* @param referenceText 参考文本
* @param vadModelId VAD模型ID
* @param asrModelId ASR模型ID
* @param llmModelId LLM模型ID
* @param ttsModelId TTS模型ID
* @param memModelId 记忆模型ID
* @param intentModelId 意图模型ID
* @param result 结果Map
*/
private void buildModuleConfig(
String assistantName,
String prompt,
String summaryMemory,
String voice,
String referenceAudio,
String referenceText,
String vadModelId,
String asrModelId,
String llmModelId,
@@ -278,8 +302,10 @@ public class ConfigServiceImpl implements ConfigService {
if (model.getConfigJson() != null) {
typeConfig.put(model.getId(), model.getConfigJson());
// 如果是TTS类型,添加private_voice属性
if ("TTS".equals(modelTypes[i]) && voice != null) {
((Map<String, Object>) model.getConfigJson()).put("private_voice", voice);
if ("TTS".equals(modelTypes[i])){
if (voice != null) ((Map<String, Object>) model.getConfigJson()).put("private_voice", voice);
if (referenceAudio != null) ((Map<String, Object>) model.getConfigJson()).put("ref_audio", referenceAudio);
if (referenceText != null) ((Map<String, Object>) model.getConfigJson()).put("ref_text", referenceText);
}
// 如果是Intent类型,且type=intent_llm,则给他添加附加模型
if ("Intent".equals(modelTypes[i])) {
@@ -25,6 +25,7 @@ import xiaozhi.common.utils.Result;
import xiaozhi.modules.device.dto.DeviceRegisterDTO;
import xiaozhi.modules.device.dto.DeviceUnBindDTO;
import xiaozhi.modules.device.dto.DeviceUpdateDTO;
import xiaozhi.modules.device.dto.DeviceManualAddDTO;
import xiaozhi.modules.device.entity.DeviceEntity;
import xiaozhi.modules.device.service.DeviceService;
import xiaozhi.modules.security.user.SecurityUser;
@@ -99,4 +100,13 @@ public class DeviceController {
deviceService.updateById(entity);
return new Result<Void>();
}
@PostMapping("/manual-add")
@Operation(summary = "手动添加设备")
@RequiresPermissions("sys:role:normal")
public Result<Void> manualAddDevice(@RequestBody @Valid DeviceManualAddDTO dto) {
UserDetail user = SecurityUser.getUser();
deviceService.manualAddDevice(user.getId(), dto);
return new Result<>();
}
}
@@ -0,0 +1,11 @@
package xiaozhi.modules.device.dto;
import lombok.Data;
@Data
public class DeviceManualAddDTO {
private String agentId;
private String board; // 设备型号
private String appVersion; // 固件版本
private String macAddress; // Mac地址
}
@@ -8,6 +8,7 @@ import xiaozhi.common.service.BaseService;
import xiaozhi.modules.device.dto.DevicePageUserDTO;
import xiaozhi.modules.device.dto.DeviceReportReqDTO;
import xiaozhi.modules.device.dto.DeviceReportRespDTO;
import xiaozhi.modules.device.dto.DeviceManualAddDTO;
import xiaozhi.modules.device.entity.DeviceEntity;
import xiaozhi.modules.device.vo.UserShowDeviceListVO;
@@ -87,5 +88,14 @@ public interface DeviceService extends BaseService<DeviceEntity> {
*/
Date getLatestLastConnectionTime(String agentId);
/**
* 手动添加设备
*/
void manualAddDevice(Long userId, DeviceManualAddDTO dto);
/**
* 更新设备连接信息
*/
void updateDeviceConnectionInfo(String agentId, String deviceId, String appVersion);
}
@@ -44,6 +44,7 @@ import xiaozhi.modules.device.vo.UserShowDeviceListVO;
import xiaozhi.modules.security.user.SecurityUser;
import xiaozhi.modules.sys.service.SysParamsService;
import xiaozhi.modules.sys.service.SysUserUtilService;
import xiaozhi.modules.device.dto.DeviceManualAddDTO;
@Slf4j
@Service
@@ -410,4 +411,30 @@ public class DeviceServiceImpl extends BaseServiceImpl<DeviceDao, DeviceEntity>
}
return 0;
}
@Override
public void manualAddDevice(Long userId, DeviceManualAddDTO dto) {
// 检查mac是否已存在
QueryWrapper<DeviceEntity> wrapper = new QueryWrapper<>();
wrapper.eq("mac_address", dto.getMacAddress());
DeviceEntity exist = baseDao.selectOne(wrapper);
if (exist != null) {
throw new RenException("该Mac地址已存在");
}
Date now = new Date();
DeviceEntity entity = new DeviceEntity();
entity.setId(dto.getMacAddress());
entity.setUserId(userId);
entity.setAgentId(dto.getAgentId());
entity.setBoard(dto.getBoard());
entity.setAppVersion(dto.getAppVersion());
entity.setMacAddress(dto.getMacAddress());
entity.setCreateDate(now);
entity.setUpdateDate(now);
entity.setLastConnectedAt(now);
entity.setCreator(userId);
entity.setUpdater(userId);
entity.setAutoUpdate(1);
baseDao.insert(entity);
}
}
@@ -104,6 +104,9 @@ public class SysParamsController {
// 验证OTA地址
validateOtaUrl(dto.getParamCode(), dto.getParamValue());
// 验证MCP地址
validateMcpUrl(dto.getParamCode(), dto.getParamValue());
sysParamsService.update(dto);
configService.getConfig(false);
return new Result<Void>();
@@ -195,4 +198,33 @@ public class SysParamsController {
throw new RenException("OTA接口验证失败:" + e.getMessage());
}
}
private void validateMcpUrl(String paramCode, String url) {
if (!paramCode.equals(Constant.SERVER_MCP_ENDPOINT)) {
return;
}
if (StringUtils.isBlank(url) || url.equals("null")) {
throw new RenException("MCP地址不能为空");
}
if (url.contains("localhost") || url.contains("127.0.0.1")) {
throw new RenException("MCP地址不能使用localhost或127.0.0.1");
}
if (!url.toLowerCase().contains("key")) {
throw new RenException("不是正确的MCP地址");
}
try {
// 发送GET请求
ResponseEntity<String> response = restTemplate.getForEntity(url, String.class);
if (response.getStatusCode() != HttpStatus.OK) {
throw new RenException("MCP接口访问失败,状态码:" + response.getStatusCode());
}
// 检查响应内容是否包含OTA相关信息
String body = response.getBody();
if (body == null || !body.contains("success")) {
throw new RenException("MCP接口返回内容格式不正确,可能不是一个真实的MCP接口");
}
} catch (Exception e) {
throw new RenException("MCP接口验证失败:" + e.getMessage());
}
}
}
@@ -26,6 +26,12 @@ public class TimbreDataDTO {
@Schema(description = "备注")
private String remark;
@Schema(description = "参考音频路径")
private String referenceAudio;
@Schema(description = "參考文本")
private String referenceText;
@Schema(description = "排序")
@Min(value = 0, message = "{sort.number}")
private long sort;
@@ -34,6 +34,12 @@ public class TimbreEntity {
@Schema(description = "备注")
private String remark;
@Schema(description = "参考音频路径")
private String referenceAudio;
@Schema(description = "參考文本")
private String referenceText;
@Schema(description = "排序")
private long sort;
@@ -25,6 +25,12 @@ public class TimbreDetailsVO implements Serializable {
@Schema(description = "备注")
private String remark;
@Schema(description = "参考音频路径")
private String referenceAudio;
@Schema(description = "參考文本")
private String referenceText;
@Schema(description = "排序")
private long sort;
@@ -0,0 +1,3 @@
ALTER TABLE `ai_tts_voice`
ADD COLUMN `reference_audio` VARCHAR(500) DEFAULT NULL COMMENT '参考音频路径' AFTER `remark`,
ADD COLUMN `reference_text` VARCHAR(500) DEFAULT NULL COMMENT '参考文本' AFTER `reference_audio`;
@@ -0,0 +1,19 @@
-- LLM意图识别配置说明
UPDATE `ai_model_config` SET
`doc_link` = NULL,
`remark` = 'LLM意图识别配置说明:
1. 使用独立的LLM进行意图识别
2. 默认使用selected_module.LLM的模型
3. 可以配置使用独立的LLM(如免费的ChatGLMLLM
4. 通用性强,但会增加处理时间
配置说明:
1. 在llm字段中指定使用的LLM模型
2. 如果不指定,则使用selected_module.LLM的模型' WHERE `id` = 'Intent_intent_llm';
-- 函数调用意图识别配置说明
UPDATE `ai_model_config` SET
`doc_link` = NULL,
`remark` = '函数调用意图识别配置说明:
1. 使用LLM的function_call功能进行意图识别
2. 需要所选择的LLM支持function_call
3. 按需调用工具,处理速度快' WHERE `id` = 'Intent_function_call';
@@ -0,0 +1,18 @@
-- 更新现有的 get_news_from_newsnow 插件配置
UPDATE ai_model_provider
SET fields = JSON_ARRAY(
JSON_OBJECT(
'key', 'url',
'type', 'string',
'label', '接口地址',
'default', 'https://newsnow.busiyi.world/api/s?id='
),
JSON_OBJECT(
'key', 'news_sources',
'type', 'string',
'label', '新闻源配置',
'default', '澎湃新闻;百度热搜;财联社'
)
)
WHERE provider_code = 'get_news_from_newsnow'
AND model_type = 'Plugin';
@@ -0,0 +1,2 @@
delete from `sys_params` where id = 113;
INSERT INTO `sys_params` (id, param_code, param_value, value_type, param_type, remark) VALUES (113, 'server.mcp_endpoint', 'null', 'string', 1, 'mcp接入点地址');
@@ -205,10 +205,38 @@ databaseChangeLog:
- sqlFile:
encoding: utf8
path: classpath:db/changelog/202506080955.sql
- changeSet:
id: 202506091720
author: shane0411
changes:
- sqlFile:
encoding: utf8
path: classpath:db/changelog/202506091720.sql
- changeSet:
id: 202506161101
author: hrz
changes:
- sqlFile:
encoding: utf8
path: classpath:db/changelog/202506161101.sql
path: classpath:db/changelog/202506161101.sql
- changeSet:
id: 202506191643
author: hrz
changes:
- sqlFile:
encoding: utf8
path: classpath:db/changelog/202506191643.sql
- changeSet:
id: 202506251620
author: Tink
changes:
- sqlFile:
encoding: utf8
path: classpath:db/changelog/202506251620.sql
- changeSet:
id: 202506261637
author: hrz
changes:
- sqlFile:
encoding: utf8
path: classpath:db/changelog/202506261637.sql
@@ -1 +1 @@
redis.call('FLUSHALL')
redis.call('FLUSHDB')
@@ -0,0 +1,85 @@
package xiaozhi.common.utils;
import static org.junit.jupiter.api.Assertions.assertEquals;
import org.junit.jupiter.api.Test;
public class AESUtilsTest {
@Test
public void testEncryptAndDecrypt() {
String key = "xiaozhi1234567890";
String plainText = "Hello, 小智!";
System.out.println("原始文本: " + plainText);
System.out.println("密钥: " + key);
// 加密
String encrypted = AESUtils.encrypt(key, plainText);
System.out.println("加密结果: " + encrypted);
// 解密
String decrypted = AESUtils.decrypt(key, encrypted);
System.out.println("解密结果: " + decrypted);
// 验证
assertEquals(plainText, decrypted, "加解密结果应该一致");
System.out.println("加解密一致性: " + plainText.equals(decrypted));
}
@Test
public void testDifferentKeyLengths() {
String[] keys = {
"1234567890123456", // 16位
"123456789012345678901234", // 24位
"12345678901234567890123456789012", // 32位
"short", // 短密钥
"verylongkeythatwillbetruncatedto32bytes" // 长密钥
};
String plainText = "测试文本";
for (String key : keys) {
String encrypted = AESUtils.encrypt(key, plainText);
String decrypted = AESUtils.decrypt(key, encrypted);
assertEquals(plainText, decrypted, "密钥长度: " + key.length());
}
}
@Test
public void testSpecialCharacters() {
String key = "xiaozhi1234567890";
String[] testTexts = {
"Hello World",
"你好世界",
"Hello, 小智!",
"特殊字符: !@#$%^&*()",
"数字123和中文混合",
"Emoji: 😀🎉🚀",
"空字符串测试",
""
};
for (String text : testTexts) {
String encrypted = AESUtils.encrypt(key, text);
String decrypted = AESUtils.decrypt(key, encrypted);
assertEquals(text, decrypted, "测试文本: " + text);
}
}
@Test
public void testCrossLanguageCompatibility() {
// 这些是Python版本生成的加密结果,用于测试跨语言兼容性
String key = "xiaozhi1234567890";
String plainText = "Hello, 小智!";
// Python版本生成的加密结果(需要运行Python测试后获取)
// String pythonEncrypted = "从Python测试中获取的加密结果";
// String decrypted = AESUtils.decrypt(key, pythonEncrypted);
// assertEquals(plainText, decrypted, "Java应该能解密Python加密的结果");
// 生成Java加密结果供Python测试
String javaEncrypted = AESUtils.encrypt(key, plainText);
System.out.println("Java加密结果供Python测试: " + javaEncrypted);
}
}
+1 -1
View File
@@ -7,6 +7,7 @@ import model from './module/model.js'
import ota from './module/ota.js'
import timbre from "./module/timbre.js"
import user from './module/user.js'
/**
* 接口地址
* 开发时自动读取使用.env.development文件
@@ -22,7 +23,6 @@ export function getServiceUrl() {
return DEV_API_SERVICE
}
/** request服务封装 */
export default {
getServiceUrl,
+30
View File
@@ -143,4 +143,34 @@ export default {
});
}).send();
},
// 获取智能体的MCP接入点地址
getAgentMcpAccessAddress(agentId, callback) {
RequestService.sendRequest()
.url(`${getServiceUrl()}/agent/mcp/address/${agentId}`)
.method('GET')
.success((res) => {
RequestService.clearRequestTime();
callback(res);
})
.networkFail(() => {
RequestService.reAjaxFun(() => {
this.getAgentMcpAccessAddress(agentId, callback);
});
}).send();
},
// 获取智能体的MCP工具列表
getAgentMcpToolsList(agentId, callback) {
RequestService.sendRequest()
.url(`${getServiceUrl()}/agent/mcp/tools/${agentId}`)
.method('GET')
.success((res) => {
RequestService.clearRequestTime();
callback(res);
})
.networkFail(() => {
RequestService.reAjaxFun(() => {
this.getAgentMcpToolsList(agentId, callback);
});
}).send();
},
}
@@ -68,4 +68,21 @@ export default {
})
}).send()
},
// 手动添加设备
manualAddDevice(params, callback) {
RequestService.sendRequest()
.url(`${getServiceUrl()}/device/manual-add`)
.method('POST')
.data(params)
.success((res) => {
RequestService.clearRequestTime();
callback(res);
})
.networkFail((err) => {
console.error('手动添加设备失败:', err);
RequestService.reAjaxFun(() => {
this.manualAddDevice(params, callback);
});
}).send();
},
}
@@ -34,6 +34,8 @@ export default {
languages: params.languageType,
name: params.voiceName,
remark: params.remark,
referenceAudio: params.referenceAudio,
referenceText: params.referenceText,
sort: params.sort,
ttsModelId: params.ttsModelId,
ttsVoice: params.voiceCode,
@@ -75,6 +77,8 @@ export default {
languages: params.languageType,
name: params.voiceName,
remark: params.remark,
referenceAudio: params.referenceAudio,
referenceText: params.referenceText,
ttsModelId: params.ttsModelId,
ttsVoice: params.voiceCode,
voiceDemo: params.voiceDemo || ''
@@ -99,6 +99,60 @@
</div>
</div>
<!-- MCP区域 -->
<div class="mcp-access-point">
<div class="mcp-container">
<!-- 左侧区域 -->
<div class="mcp-left">
<div class="mcp-header">
<h3 class="bold-title">MCP接入点</h3>
</div>
<div class="url-header">
<div class="address-desc">
<span>以下是智能体的MCP接入点地址</span>
<a href="https://github.com/xinnan-tech/xiaozhi-esp32-server/blob/main/docs/mcp-endpoint-integration.md"
target="_blank" class="doc-link">查看接入点使用文档</a>
</div>
</div>
<el-input v-model="mcpUrl" readonly class="url-input">
<template #suffix>
<el-button @click="copyUrl" class="inner-copy-btn" icon="el-icon-document-copy">
复制
</el-button>
</template>
</el-input>
</div>
<!-- 右侧区域 -->
<div class="mcp-right">
<div class="mcp-header">
<h3 class="bold-title">接入点状态</h3>
</div>
<div class="status-container">
<span class="status-indicator" :class="mcpStatus"></span>
<span class="status-text">{{
mcpStatus === 'connected' ? '已连接' :
mcpStatus === 'loading' ? '加载中...' : '未连接'
}}</span>
<button class="refresh-btn" @click="refreshStatus">
<span class="refresh-icon"></span>
<span>刷新</span>
</button>
</div>
<div class="mcp-tools-list">
<div v-if="mcpTools.length > 0" class="tools-grid">
<el-button v-for="tool in mcpTools" :key="tool" size="small" class="tool-btn" plain>
{{ tool }}
</el-button>
</div>
<div v-else class="no-tools">
<span>暂无可用工具</span>
</div>
</div>
</div>
</div>
</div>
<div class="drawer-footer">
<el-button @click="closeDialog">取消</el-button>
<el-button type="primary" @click="saveSelection">保存配置</el-button>
@@ -107,6 +161,8 @@
</template>
<script>
import Api from '@/apis/api';
export default {
props: {
value: Boolean,
@@ -117,6 +173,10 @@ export default {
allFunctions: {
type: Array,
default: () => []
},
agentId: {
type: String,
required: true
}
},
data() {
@@ -134,6 +194,10 @@ export default {
// 添加一个标志位来跟踪是否已经保存
hasSaved: false,
loading: false,
mcpUrl: "",
mcpStatus: "disconnected",
mcpTools: [],
}
},
computed: {
@@ -177,6 +241,10 @@ export default {
});
// 右侧默认指向第一个
this.currentFunction = this.selectedList[0] || null;
// 加载MCP数据
this.loadMcpAddress();
this.loadMcpTools();
}
},
dialogVisible(newVal) {
@@ -184,6 +252,60 @@ export default {
}
},
methods: {
copyUrl() {
const textarea = document.createElement('textarea');
textarea.value = this.mcpUrl;
textarea.style.position = 'fixed'; // 防止页面滚动
document.body.appendChild(textarea);
textarea.select();
try {
const successful = document.execCommand('copy');
if (successful) {
this.$message.success('已复制到剪贴板');
} else {
this.$message.error('复制失败,请手动复制');
}
} catch (err) {
this.$message.error('复制失败,请手动复制');
console.error('复制失败:', err);
} finally {
document.body.removeChild(textarea);
}
},
refreshStatus() {
this.mcpStatus = "loading";
this.loadMcpTools();
},
// 加载MCP接入点地址
loadMcpAddress() {
Api.agent.getAgentMcpAccessAddress(this.agentId, (res) => {
if (res.data.code === 0) {
this.mcpUrl = res.data.data || "";
} else {
this.mcpUrl = "";
console.error('获取MCP地址失败:', res.data.msg);
}
});
},
// 加载MCP工具列表
loadMcpTools() {
Api.agent.getAgentMcpToolsList(this.agentId, (res) => {
if (res.data.code === 0) {
this.mcpTools = res.data.data || [];
// 根据工具列表更新状态
this.mcpStatus = this.mcpTools.length > 0 ? "connected" : "disconnected";
} else {
this.mcpTools = [];
this.mcpStatus = "disconnected";
console.error('获取MCP工具列表失败:', res.data.msg);
}
});
},
flushArray(key) {
const text = this.textCache[key] || '';
const arr = text
@@ -298,7 +420,7 @@ export default {
display: grid;
grid-template-columns: max-content max-content 1fr;
gap: 12px;
height: calc(70vh - 60px);
height: calc(58vh);
}
.custom-header {
@@ -514,4 +636,202 @@ export default {
::v-deep .el-checkbox__label {
display: none;
}
.mcp-access-point {
border-top: 1px solid #EBEEF5;
padding: 20px 24px;
text-align: left;
}
.mcp-header {
.bold-title {
font-size: 18px;
font-weight: bold;
margin: 5px 0 30px 0;
}
}
.mcp-container {
display: flex;
justify-content: space-between;
gap: 30px;
}
.mcp-left,
.mcp-right {
flex: 1;
}
.url-header {
margin-bottom: 8px;
color: black;
h4 {
margin: 0 0 15px 0;
font-size: 16px;
font-weight: normal;
}
.address-desc {
display: flex;
align-items: center;
font-size: 14px;
margin-bottom: 12px;
.doc-link {
color: #1677ff;
text-decoration: none;
margin-left: 4px;
&:hover {
text-decoration: underline;
}
}
}
}
.url-input {
border-radius: 4px 0 0 4px;
font-size: 14px;
height: 36px;
box-sizing: border-box;
background-color: #f5f5f5;
}
::v-deep .el-input__inner {
background-color: #f5f5f5;
padding-right: 80px;
}
.url-input {
::v-deep .el-input__suffix {
right: 0;
display: flex;
align-items: center;
padding-right: 10px;
.inner-copy-btn {
pointer-events: auto;
border: none;
background: #1677ff;
color: white;
padding: 6px;
margin-top: 4px;
margin-left: 4px;
}
}
}
.mcp-right {
h4 {
margin: 0 0 10px 0;
font-size: 16px;
font-weight: normal;
color: black;
}
}
.status-container {
display: flex;
align-items: center;
.status-indicator {
display: inline-block;
width: 8px;
height: 8px;
border-radius: 50%;
margin-right: 8px;
&.disconnected {
background-color: #909399;
/* 灰色 - 未连接 */
}
&.connected {
background-color: #67C23A;
/* 绿色 - 已连接 */
}
&.loading {
background-color: #E6A23C;
/* 橙色 - 加载中 */
animation: pulse 1.5s infinite;
}
}
.status-text {
font-size: 14px;
margin-right: 10px;
}
.refresh-btn {
display: flex;
align-items: center;
padding: 2px 10px;
background: white;
color: black;
border: 1px solid #DCDFE6;
border-radius: 4px;
cursor: pointer;
font-size: 14px;
transition: all 0.3s;
&:hover {
background: #1677ff;
color: white;
border-color: #1677ff;
}
.refresh-icon {
margin-right: 6px;
font-size: 14px;
}
}
}
@keyframes pulse {
0% {
opacity: 1;
}
50% {
opacity: 0.4;
}
100% {
opacity: 1;
}
}
.mcp-tools-list {
margin-top: 10px;
.tools-grid {
display: flex;
flex-wrap: wrap;
gap: 8px;
}
.tool-btn {
padding: 6px 12px;
border-color: #1677ff;
color: #1677ff;
background-color: white;
font-size: 12px;
&:hover {
background-color: #1677ff;
color: white;
border-color: #1677ff;
}
}
.no-tools {
text-align: center;
color: #909399;
font-size: 14px;
padding: 10px 0;
}
}
</style>
@@ -0,0 +1,158 @@
<template>
<el-dialog title="手动添加设备" :visible="visible" @close="handleClose" width="30%" center>
<div class="dialog-content">
<el-form :model="deviceForm" :rules="rules" ref="deviceForm" label-width="100px">
<el-form-item label="设备型号" prop="board">
<el-select v-model="deviceForm.board" placeholder="请选择设备型号" style="width: 100%">
<el-option
v-for="item in firmwareTypes"
:key="item.key"
:label="item.name"
:value="item.key">
</el-option>
</el-select>
</el-form-item>
<el-form-item label="固件版本" prop="appVersion">
<el-input v-model="deviceForm.appVersion" placeholder="请输入固件版本"></el-input>
</el-form-item>
<el-form-item label="Mac地址" prop="macAddress">
<el-input v-model="deviceForm.macAddress" placeholder="请输入Mac地址"></el-input>
</el-form-item>
</el-form>
</div>
<div style="display: flex;margin: 15px 15px;gap: 7px;">
<div class="dialog-btn" @click="submitForm">确定</div>
<div class="dialog-btn cancel-btn" @click="cancel">取消</div>
</div>
</el-dialog>
</template>
<script>
import Api from '@/apis/api';
export default {
name: 'ManualAddDeviceDialog',
props: {
visible: { type: Boolean, required: true },
agentId: { type: String, required: true }
},
data() {
// MAC地址验证规则
const validateMac = (rule, value, callback) => {
const macRegex = /^([0-9A-Fa-f]{2}[:-]){5}([0-9A-Fa-f]{2})$/;
if (!value) {
callback(new Error('请输入Mac地址'));
} else if (!macRegex.test(value)) {
callback(new Error('请输入正确的Mac地址格式,例如:00:1A:2B:3C:4D:5E'));
} else {
callback();
}
};
return {
deviceForm: {
board: '',
appVersion: '',
macAddress: ''
},
firmwareTypes: [],
rules: {
board: [
{ required: true, message: '请选择设备型号', trigger: 'change' }
],
appVersion: [
{ required: true, message: '请输入固件版本', trigger: 'blur' }
],
macAddress: [
{ required: true, validator: validateMac, trigger: 'blur' }
]
}
}
},
created() {
this.getFirmwareTypes();
},
methods: {
async getFirmwareTypes() {
try {
const res = await Api.dict.getDictDataByType('FIRMWARE_TYPE');
this.firmwareTypes = res.data;
} catch (error) {
console.error('获取固件类型失败:', error);
this.$message.error(error.message || '获取固件类型失败');
}
},
submitForm() {
this.$refs.deviceForm.validate((valid) => {
if (valid) {
this.addDevice();
}
});
},
addDevice() {
const params = {
agentId: this.agentId,
...this.deviceForm
};
Api.device.manualAddDevice(params, ({ data }) => {
if (data.code === 0) {
this.$message.success('设备添加成功');
this.$emit('refresh');
this.closeDialog();
} else {
this.$message.error(data.msg || '添加失败');
}
});
},
closeDialog() {
this.$emit('update:visible', false);
this.$refs.deviceForm.resetFields();
},
cancel() {
this.closeDialog();
},
handleClose() {
this.closeDialog();
}
}
}
</script>
<style scoped>
.dialog-content {
padding: 0 20px;
}
.dialog-btn {
cursor: pointer;
flex: 1;
border-radius: 23px;
background: #5778ff;
height: 40px;
font-weight: 500;
font-size: 12px;
color: #fff;
line-height: 40px;
text-align: center;
}
.cancel-btn {
background: #e6ebff;
border: 1px solid #adbdff;
color: #5778ff;
}
::v-deep .el-dialog {
border-radius: 15px;
box-shadow: 0 4px 12px rgba(0, 0, 0, 0.1);
}
::v-deep .el-dialog__body {
padding: 20px 6px;
}
::v-deep .el-form-item {
margin-bottom: 20px;
}
</style>
+118 -106
View File
@@ -1,79 +1,66 @@
<template>
<el-dialog
:visible.sync="localVisible"
width="85%"
@close="handleClose"
:show-close="false"
:append-to-body="true"
:close-on-click-modal="true">
<el-dialog :visible.sync="localVisible" width="90%" @close="handleClose" :show-close="false" :append-to-body="true"
:close-on-click-modal="true">
<button class="custom-close-btn" @click="handleClose">
×
</button>
<div class="scroll-wrapper">
<div
class="table-container"
ref="tableContainer"
@scroll="handleScroll">
<el-table
v-loading="loading"
:data="filteredTtsModels"
style="width: 100%;"
class="data-table"
header-row-class-name="table-header"
:fit="true"
element-loading-text="拼命加载中"
element-loading-spinner="el-icon-loading"
element-loading-background="rgba(0, 0, 0, 0.8)">
<div class="table-container" ref="tableContainer" @scroll="handleScroll">
<el-table v-loading="loading" :data="filteredTtsModels" style="width: 100%;" class="data-table"
header-row-class-name="table-header" :fit="true" element-loading-text="拼命加载中"
element-loading-spinner="el-icon-loading" element-loading-background="rgba(0, 0, 0, 0.8)">
<el-table-column label="选择" width="50" align="center">
<template slot-scope="scope">
<el-checkbox v-model="scope.row.selected"></el-checkbox>
</template>
</el-table-column>
<el-table-column label="音色编码" align="center" min-width="50">
<el-table-column label="音色编码" align="center">
<template slot-scope="scope">
<el-input v-if="scope.row.editing" v-model="scope.row.voiceCode"></el-input>
<span v-else>{{ scope.row.voiceCode }}</span>
</template>
</el-table-column>
<el-table-column label="音色名称" align="center" min-width="50">
<el-table-column label="音色名称" align="center">
<template slot-scope="scope">
<el-input v-if="scope.row.editing" v-model="scope.row.voiceName"></el-input>
<span v-else>{{ scope.row.voiceName }}</span>
</template>
</el-table-column>
<el-table-column label="语言类型" align="center" min-width="50">
<el-table-column label="语言类型" align="center">
<template slot-scope="scope">
<el-input v-if="scope.row.editing" v-model="scope.row.languageType"></el-input>
<span v-else>{{ scope.row.languageType }}</span>
</template>
</el-table-column>
<el-table-column label="试听" align="center" min-width="100px" class-name="audio-column">
<el-table-column v-if="!showReferenceColumns" label="试听" align="center" class-name="audio-column">
<template slot-scope="scope">
<div class="custom-audio-container">
<el-input
v-if="scope.row.editing"
v-model="scope.row.voiceDemo"
placeholder="请输入MP3地址"
size="mini"
class="audio-input">
<el-input v-if="scope.row.editing" v-model="scope.row.voiceDemo" placeholder="请输入MP3地址"
class="audio-input">
</el-input>
<AudioPlayer v-else-if="isValidAudioUrl(scope.row.voiceDemo)" :audioUrl="scope.row.voiceDemo"/>
<AudioPlayer v-else-if="isValidAudioUrl(scope.row.voiceDemo)" :audioUrl="scope.row.voiceDemo" />
</div>
</template>
</el-table-column>
<el-table-column label="备注" align="center">
<el-table-column v-if="!showReferenceColumns" label="备注" align="center">
<template slot-scope="scope">
<el-input
v-if="scope.row.editing"
type="textarea"
:rows="1"
autosize
v-model="scope.row.remark"
placeholder="这里是备注"
class="remark-input"></el-input>
<el-input v-if="scope.row.editing" type="textarea" :rows="1" autosize v-model="scope.row.remark"
placeholder="这里是备注" class="remark-input"></el-input>
<span v-else>{{ scope.row.remark }}</span>
</template>
</el-table-column>
<el-table-column v-if="showReferenceColumns" label="克隆音频路径" align="center">
<template slot-scope="scope">
<el-input v-if="scope.row.editing" v-model="scope.row.referenceAudio" placeholder="这里是克隆音频路径"></el-input>
<span v-else>{{ scope.row.referenceAudio }}</span>
</template>
</el-table-column>
<el-table-column v-if="showReferenceColumns" label="克隆音频文本" align="center">
<template slot-scope="scope">
<el-input v-if="scope.row.editing" v-model="scope.row.referenceText" placeholder="这里是克隆音频对应文本"></el-input>
<span v-else>{{ scope.row.referenceText }}</span>
</template>
</el-table-column>
<el-table-column label="操作" align="center" width="150">
<template slot-scope="scope">
<template v-if="!scope.row.editing">
@@ -84,12 +71,7 @@
删除
</el-button>
</template>
<el-button
v-else
type="success"
size="mini"
@click="saveEdit(scope.row)"
class="save-Tts">保存
<el-button v-else type="success" size="mini" @click="saveEdit(scope.row)" class="save-Tts">保存
</el-button>
</template>
</el-table-column>
@@ -110,21 +92,19 @@
<el-button type="primary" size="mini" @click="addNew" style="background: #5bc98c;border: None;">
新增
</el-button>
<el-button type="primary"
size="mini"
@click="deleteRow(filteredTtsModels.filter(row => row.selected))"
style="background: red;border:None">删除
<el-button type="primary" size="mini" @click="deleteRow(filteredTtsModels.filter(row => row.selected))"
style="background: red;border:None">删除
</el-button>
</div>
</el-dialog>
</template>
<script>
import AudioPlayer from './AudioPlayer.vue'
import Api from "@/apis/api";
import AudioPlayer from './AudioPlayer.vue';
export default {
components: {AudioPlayer},
components: { AudioPlayer },
props: {
visible: {
type: Boolean,
@@ -133,6 +113,10 @@ export default {
ttsModelId: {
type: String,
required: true
},
modelConfig: {
type: Object,
default: null
}
},
data() {
@@ -151,6 +135,7 @@ export default {
selectAll: false,
selectedRows: [],
loading: false,
showReferenceColumns: false, // 控制是否显示参考列
};
},
watch: {
@@ -158,12 +143,19 @@ export default {
this.localVisible = newVal;
if (newVal) {
this.currentPage = 1;
this.updateShowReferenceColumns(); // 更新显示状态
this.loadData(); // 对话框显示时加载数据
this.$nextTick(() => {
this.updateScrollbar();
});
}
},
modelConfig: {
handler(newVal) {
this.updateShowReferenceColumns();
},
immediate: true
},
filteredTtsModels() {
this.$nextTick(() => {
this.updateScrollbar();
@@ -173,7 +165,7 @@ export default {
computed: {
filteredTtsModels() {
return this.ttsModels.filter(model =>
model.voiceName.toLowerCase().includes(this.searchQuery.toLowerCase())
model.voiceName.toLowerCase().includes(this.searchQuery.toLowerCase())
);
}
},
@@ -189,6 +181,16 @@ export default {
window.removeEventListener('mousemove', this.handleDrag);
},
methods: {
// 更新是否显示参考列
updateShowReferenceColumns() {
if (this.modelConfig && this.modelConfig.configJson) {
const providerType = this.modelConfig.configJson.type;
this.showReferenceColumns = ['fishspeech', 'gpt_sovits_v2', 'gpt_sovits_v3'].includes(providerType);
} else {
this.showReferenceColumns = false;
}
},
loadData() {
this.loading = true;
const params = {
@@ -206,6 +208,8 @@ export default {
voiceName: item.name || '未命名音色',
languageType: item.languages || '',
remark: item.remark || '',
referenceAudio: item.referenceAudio || '',
referenceText: item.referenceText || '',
voiceDemo: item.voiceDemo || '',
selected: false,
editing: false,
@@ -237,6 +241,7 @@ export default {
this.total = 0;
this.selectAll = false;
this.searchQuery = '';
this.showReferenceColumns = false;
this.localVisible = false;
this.$emit('update:visible', false);
@@ -249,7 +254,7 @@ export default {
if (!container || !scrollbarThumb || !scrollbarTrack) return;
const {scrollHeight, clientHeight} = container;
const { scrollHeight, clientHeight } = container;
const trackHeight = scrollbarTrack.clientHeight;
const thumbHeight = Math.max((clientHeight / scrollHeight) * trackHeight, 20);
@@ -264,7 +269,7 @@ export default {
if (!container || !scrollbarThumb || !scrollbarTrack) return;
const {scrollHeight, clientHeight, scrollTop} = container;
const { scrollHeight, clientHeight, scrollTop } = container;
const trackHeight = scrollbarTrack.clientHeight;
const thumbHeight = scrollbarThumb.clientHeight;
const maxTop = trackHeight - thumbHeight;
@@ -332,7 +337,7 @@ export default {
startEdit(row) {
row.editing = true;
this.$set(row, 'originalData', {...row});
this.$set(row, 'originalData', { ...row });
},
saveEdit(row) {
@@ -356,6 +361,12 @@ export default {
sort: row.sort
};
// 只有在显示参考列的情况下才添加参考字段
if (this.showReferenceColumns) {
params.referenceAudio = row.referenceAudio;
params.referenceText = row.referenceText;
}
let res;
if (row.id) {
// 已有ID,执行更新操作
@@ -423,8 +434,8 @@ export default {
}
const maxSort = this.ttsModels.length > 0
? Math.max(...this.ttsModels.map(item => Number(item.sort) || 0))
: 0;
? Math.max(...this.ttsModels.map(item => Number(item.sort) || 0))
: 0;
const newRow = {
voiceCode: '',
@@ -432,6 +443,8 @@ export default {
languageType: '中文',
voiceDemo: '',
remark: '',
referenceAudio: '',
referenceText: '',
selected: false,
editing: true,
sort: maxSort + 1
@@ -441,56 +454,56 @@ export default {
},
deleteRow(row) {
// 处理单个音色或音色数组
const voices = Array.isArray(row) ? row : [row];
// 处理单个音色或音色数组
const voices = Array.isArray(row) ? row : [row];
if (Array.isArray(row) && row.length === 0) {
this.$message.warning("请先选择需要删除的音色");
return;
if (Array.isArray(row) && row.length === 0) {
this.$message.warning("请先选择需要删除的音色");
return;
}
const voiceCount = voices.length;
this.$confirm(`确定要删除选中的${voiceCount}个音色吗?`, "警告", {
confirmButtonText: "确定",
cancelButtonText: "取消",
type: "warning",
distinguishCancelAndClose: true
}).then(() => {
const ids = voices.map(voice => voice.id);
if (ids.some(id => !id)) {
this.$message.error("存在无效的音色ID");
return;
}
const voiceCount = voices.length;
this.$confirm(`确定要删除选中的${voiceCount}个音色吗?`, "警告", {
confirmButtonText: "确定",
cancelButtonText: "取消",
type: "warning",
distinguishCancelAndClose: true
}).then(() => {
const ids = voices.map(voice => voice.id);
if (ids.some(id => !id)) {
this.$message.error("存在无效的音色ID");
return;
}
Api.timbre.deleteVoice(ids, ({data}) => {
if (data.code === 0) {
this.$message.success({
message: `成功删除${voiceCount}个参数`,
showClose: true
});
this.loadData(); // 刷新参数列表
} else {
this.$message.error({
message: data.msg || '删除失败,请重试',
showClose: true
});
}
Api.timbre.deleteVoice(ids, ({ data }) => {
if (data.code === 0) {
this.$message.success({
message: `成功删除${voiceCount}个参数`,
showClose: true
});
this.loadData(); // 刷新参数列表
} else {
this.$message.error({
message: data.msg || '删除失败,请重试',
showClose: true
});
}
});
}).catch(action => {
if (action === 'cancel') {
this.$message({
type: 'info',
message: '已取消删除操作',
duration: 1000
});
} else {
this.$message({
type: 'info',
message: '操作已关闭',
duration: 1000
});
}
if (action === 'cancel') {
this.$message({
type: 'info',
message: '已取消删除操作',
duration: 1000
});
} else {
this.$message({
type: 'info',
message: '操作已关闭',
duration: 1000
});
}
});
},
@@ -502,7 +515,6 @@ export default {
</script>
<style lang="scss" scoped>
::v-deep .el-dialog {
border-radius: 8px !important;
overflow: hidden;
@@ -695,6 +707,7 @@ export default {
white-space: pre-wrap !important;
word-break: break-all !important;
}
/* 按钮组定位调整 */
.action-buttons {
position: static;
@@ -719,5 +732,4 @@ export default {
flex-shrink: 0;
margin: 2px !important;
}
</style>
@@ -77,7 +77,10 @@
{{ isAllSelected ? '取消全选' : '全选' }}
</el-button>
<el-button type="success" size="mini" class="add-device-btn" @click="handleAddDevice">
新增
验证码绑定
</el-button>
<el-button type="success" size="mini" class="add-device-btn" @click="handleManualAddDevice">
手动添加
</el-button>
<el-button size="mini" type="danger" icon="el-icon-delete" @click="deleteSelected">解绑</el-button>
</div>
@@ -103,6 +106,8 @@
<AddDeviceDialog :visible.sync="addDeviceDialogVisible" :agent-id="currentAgentId"
@refresh="fetchBindDevices(currentAgentId)"/>
<ManualAddDeviceDialog :visible.sync="manualAddDeviceDialogVisible" :agent-id="currentAgentId"
@refresh="fetchBindDevices(currentAgentId)"/>
</div>
</template>
@@ -110,13 +115,19 @@
<script>
import Api from '@/apis/api';
import AddDeviceDialog from "@/components/AddDeviceDialog.vue";
import ManualAddDeviceDialog from "@/components/ManualAddDeviceDialog.vue";
import HeaderBar from "@/components/HeaderBar.vue";
export default {
components: {HeaderBar, AddDeviceDialog},
components: {
HeaderBar,
AddDeviceDialog,
ManualAddDeviceDialog
},
data() {
return {
addDeviceDialogVisible: false,
manualAddDeviceDialogVisible: false,
selectedDevices: [],
isAllSelected: false,
searchKeyword: "",
@@ -252,6 +263,9 @@ export default {
handleAddDevice() {
this.addDeviceDialogVisible = true;
},
handleManualAddDevice() {
this.manualAddDeviceDialogVisible = true;
},
submitRemark(row) {
if (row._submitting) return;
+4 -2
View File
@@ -128,7 +128,7 @@
<ModelEditDialog :modelType="activeTab" :visible.sync="editDialogVisible" :modelData="editModelData"
@save="handleModelSave" />
<TtsModel :visible.sync="ttsDialogVisible" :ttsModelId="selectedTtsModelId" />
<TtsModel :visible.sync="ttsDialogVisible" :ttsModelId="selectedTtsModelId" :modelConfig="selectedModelConfig" />
<AddModelDialog :modelType="activeTab" :visible.sync="addDialogVisible" @confirm="handleAddConfirm" />
</div>
<el-footer>
@@ -162,7 +162,8 @@ export default {
total: 0,
selectedModels: [],
isAllSelected: false,
loading: false
loading: false,
selectedModelConfig: {}
};
},
@@ -211,6 +212,7 @@ export default {
},
openTtsDialog(row) {
this.selectedTtsModelId = row.id;
this.selectedModelConfig = row;
this.ttsDialogVisible = true;
},
headerCellClassName({ column, columnIndex }) {
+1 -1
View File
@@ -132,7 +132,7 @@
</div>
<function-dialog v-model="showFunctionDialog" :functions="currentFunctions" :all-functions="allFunctions"
@update-functions="handleUpdateFunctions" @dialog-closed="handleDialogClosed" />
:agent-id="$route.query.agentId" @update-functions="handleUpdateFunctions" @dialog-closed="handleDialogClosed" />
</div>
</template>
+9 -6
View File
@@ -104,7 +104,8 @@ wakeup_words:
- "喵喵同学"
- "小滨小滨"
- "小冰小冰"
# MCP接入点地址
mcp_endpoint: 你的接入点 websocket地址
# 插件的基础配置
plugins:
# 获取天气插件的配置,这里填写你的api_key
@@ -120,7 +121,9 @@ plugins:
society_rss_url: "https://www.chinanews.com.cn/rss/society.xml"
world_rss_url: "https://www.chinanews.com.cn/rss/world.xml"
finance_rss_url: "https://www.chinanews.com.cn/rss/finance.xml"
get_news_from_newsnow: {"url": "https://newsnow.busiyi.world/api/s?id="}
get_news_from_newsnow:
url: "https://newsnow.busiyi.world/api/s?id="
news_sources: "澎湃新闻;百度热搜;财联社"
home_assistant:
devices:
- 客厅,玩具灯,switch.cuco_cn_460494544_cp1_on_p_2_1
@@ -158,7 +161,7 @@ end_prompt:
enable: true # 是否开启结束语
# 结束语
prompt: |
请你以时间过得真快未来头,用富有感情、依依不舍的话来结束这场对话吧!
请你以"时间过得真快"未来头,用富有感情、依依不舍的话来结束这场对话吧!
# 具体处理时选择的模块(The module selected for specific processing)
selected_module:
@@ -195,7 +198,7 @@ Intent:
# 如果你的不想使用selected_module.LLM意图识别,这里最好使用独立的LLM作为意图识别,例如使用免费的ChatGLMLLM
llm: ChatGLMLLM
# plugins_func/functions下的模块,可以通过配置,选择加载哪个模块,加载后对话支持相应的function调用
# 系统默认已经记载handle_exit_intent(退出识别)”、“play_music(音乐播放)插件,请勿重复加载
# 系统默认已经记载"handle_exit_intent(退出识别)"、"play_music(音乐播放)"插件,请勿重复加载
# 下面是加载查天气、角色切换、加载查新闻的插件示例
functions:
- get_weather
@@ -205,7 +208,7 @@ Intent:
# 不需要动type
type: function_call
# plugins_func/functions下的模块,可以通过配置,选择加载哪个模块,加载后对话支持相应的function调用
# 系统默认已经记载handle_exit_intent(退出识别)”、“play_music(音乐播放)插件,请勿重复加载
# 系统默认已经记载"handle_exit_intent(退出识别)"、"play_music(音乐播放)"插件,请勿重复加载
# 下面是加载查天气、角色切换、加载查新闻的插件示例
functions:
- change_role
@@ -765,4 +768,4 @@ TTS:
# 支持声音克隆,可自行上传音频,填入voice参数,voice参数为空时,使用默认声音
access_token: "U4YdYXVfpwWnk2t5Gp822zWPCuORyeJL"
voice: "OUeAo1mhq6IBExi"
output_dir: tmp/
output_dir: tmp/
+5 -5
View File
@@ -5,7 +5,7 @@ from config.config_loader import load_config
from config.settings import check_config_file
from datetime import datetime
SERVER_VERSION = "0.5.7"
SERVER_VERSION = "0.6.1"
_logger_initialized = False
@@ -88,7 +88,7 @@ def setup_logging():
# 输出到文件 - 统一目录,按大小轮转
# 日志文件完整路径
log_file_path = os.path.join(log_dir, log_file)
# 添加日志处理器
logger.add(
log_file_path,
@@ -146,14 +146,14 @@ def update_module_string(selected_module_str):
level=log_config.get("log_level", "INFO"),
filter=formatter,
)
# 更新文件日志配置 - 统一目录,按大小轮转
log_dir = log_config.get("log_dir", "tmp")
log_file = log_config.get("log_file", "server.log")
# 日志文件完整路径
log_file_path = os.path.join(log_dir, log_file)
logger.add(
log_file_path,
format=log_format_file,
@@ -8,6 +8,7 @@ from config.config_loader import get_private_config_from_api
from core.utils.auth import AuthToken
import base64
from typing import Tuple, Optional
from plugins_func.register import Action
TAG = __name__
@@ -122,7 +123,8 @@ class VisionHandler:
return_json = {
"success": True,
"result": result,
"action": Action.RESPONSE.name,
"response": result,
}
response = web.Response(
+38 -112
View File
@@ -10,7 +10,6 @@ import threading
import traceback
import subprocess
import websockets
from core.handle.mcpHandle import call_mcp_tool
from core.utils.util import (
extract_json_from_string,
check_vad_update,
@@ -18,7 +17,6 @@ from core.utils.util import (
filter_sensitive_info,
)
from typing import Dict, Any
from core.mcp.manager import MCPManager
from core.utils.modules_initialize import (
initialize_modules,
initialize_tts,
@@ -30,7 +28,7 @@ from concurrent.futures import ThreadPoolExecutor
from core.utils.dialogue import Message, Dialogue
from core.providers.asr.dto.dto import InterfaceType
from core.handle.textHandle import handleTextMessage
from core.handle.functionHandler import FunctionHandler
from core.providers.tools.unified_tool_handler import UnifiedToolHandler
from plugins_func.loadplugins import auto_import_modules
from plugins_func.register import Action, ActionResponse
from core.auth import AuthMiddleware, AuthenticationError
@@ -112,8 +110,7 @@ class ConnectionHandler:
# vad相关变量
self.client_audio_buffer = bytearray()
self.client_have_voice = False
self.client_have_voice_last_time = 0.0
self.client_no_voice_last_time = 0.0
self.last_activity_time = 0.0 # 统一的活动时间戳(毫秒)
self.client_voice_stop = False
# asr相关变量
@@ -144,10 +141,10 @@ class ConnectionHandler:
self.load_function_plugin = False
self.intent_type = "nointent"
self.timeout_task = None
self.timeout_seconds = (
int(self.config.get("close_connection_no_voice_time", 120)) + 60
) # 在原来第一道关闭的基础上加60秒,进行二道关闭
self.timeout_task = None
# {"mcp":true} 表示启用MCP功能
self.features = None
@@ -188,6 +185,9 @@ class ConnectionHandler:
self.websocket = ws
self.device_id = self.headers.get("device-id", None)
# 初始化活动时间戳
self.last_activity_time = time.time() * 1000
# 启动超时检查任务
self.timeout_task = asyncio.create_task(self._check_timeout())
@@ -242,18 +242,10 @@ class ConnectionHandler:
# 立即关闭连接,不等待记忆保存完成
await self.close(ws)
async def reset_timeout(self):
"""重置超时计时器"""
if self.timeout_task:
self.timeout_task.cancel()
self.timeout_task = asyncio.create_task(self._check_timeout())
async def _route_message(self, message):
"""消息路由"""
# 重置超时计时器
await self.reset_timeout()
if isinstance(message, str):
self.last_activity_time = time.time() * 1000
await handleTextMessage(self, message)
elif isinstance(message, bytes):
if self.vad is None:
@@ -474,6 +466,8 @@ class ConnectionHandler:
self.max_output_size = int(private_config["device_max_output_size"])
if private_config.get("chat_history_conf", None) is not None:
self.chat_history_conf = int(private_config["chat_history_conf"])
if private_config.get("mcp_endpoint", None) is not None:
self.config["mcp_endpoint"] = private_config["mcp_endpoint"]
try:
modules = initialize_modules(
self.logger,
@@ -585,14 +579,12 @@ class ConnectionHandler:
self.intent.set_llm(self.llm)
self.logger.bind(tag=TAG).info("使用主LLM作为意图识别模型")
"""加载插件"""
self.func_handler = FunctionHandler(self)
self.mcp_manager = MCPManager(self)
"""加载统一工具处理器"""
self.func_handler = UnifiedToolHandler(self)
"""加载MCP工具"""
asyncio.run_coroutine_threadsafe(
self.mcp_manager.initialize_servers(), self.loop
)
# 异步初始化工具处理器
if hasattr(self, "loop") and self.loop:
asyncio.run_coroutine_threadsafe(self.func_handler._initialize(), self.loop)
def change_system_prompt(self, prompt):
self.prompt = prompt
@@ -610,12 +602,6 @@ class ConnectionHandler:
functions = None
if self.intent_type == "function_call" and hasattr(self, "func_handler"):
functions = self.func_handler.get_functions()
if hasattr(self, "mcp_client"):
mcp_tools = self.mcp_client.get_available_tools()
if mcp_tools is not None and len(mcp_tools) > 0:
if functions is None:
functions = []
functions.extend(mcp_tools)
response_message = []
try:
@@ -627,8 +613,7 @@ class ConnectionHandler:
)
memory_str = future.result()
uuid_str = str(uuid.uuid4()).replace("-", "")
self.sentence_id = uuid_str
self.sentence_id = str(uuid.uuid4().hex)
if self.intent_type == "function_call" and functions is not None:
# 使用支持functions的streaming接口
@@ -733,37 +718,13 @@ class ConnectionHandler:
"arguments": function_arguments,
}
# 处理Server端MCP工具调用
if self.mcp_manager.is_mcp_tool(function_name):
result = self._handle_mcp_tool_call(function_call_data)
elif hasattr(self, "mcp_client") and self.mcp_client.has_tool(
function_name
):
# 如果是小智端MCP工具调用
self.logger.bind(tag=TAG).debug(
f"调用小智端MCP工具: {function_name}, 参数: {function_arguments}"
)
try:
result = asyncio.run_coroutine_threadsafe(
call_mcp_tool(
self, self.mcp_client, function_name, function_arguments
),
self.loop,
).result()
self.logger.bind(tag=TAG).debug(f"MCP工具调用结果: {result}")
result = ActionResponse(
action=Action.REQLLM, result=result, response=""
)
except Exception as e:
self.logger.bind(tag=TAG).error(f"MCP工具调用失败: {e}")
result = ActionResponse(
action=Action.REQLLM, result="MCP工具调用失败", response=""
)
else:
# 处理系统函数
result = self.func_handler.handle_llm_function_call(
# 使用统一工具处理器处理所有工具调用
result = asyncio.run_coroutine_threadsafe(
self.func_handler.handle_llm_function_call(
self, function_call_data
)
),
self.loop,
).result()
self._handle_function_result(result, function_call_data)
# 存储对话内容
@@ -786,48 +747,6 @@ class ConnectionHandler:
return True
def _handle_mcp_tool_call(self, function_call_data):
function_arguments = function_call_data["arguments"]
function_name = function_call_data["name"]
try:
args_dict = function_arguments
if isinstance(function_arguments, str):
try:
args_dict = json.loads(function_arguments)
except json.JSONDecodeError:
self.logger.bind(tag=TAG).error(
f"无法解析 function_arguments: {function_arguments}"
)
return ActionResponse(
action=Action.REQLLM, result="参数解析失败", response=""
)
tool_result = asyncio.run_coroutine_threadsafe(
self.mcp_manager.execute_tool(function_name, args_dict), self.loop
).result()
# meta=None content=[TextContent(type='text', text='北京当前天气:\n温度: 21°C\n天气: 晴\n湿度: 6%\n风向: 西北 风\n风力等级: 5级', annotations=None)] isError=False
content_text = ""
if tool_result is not None and tool_result.content is not None:
for content in tool_result.content:
content_type = content.type
if content_type == "text":
content_text = content.text
elif content_type == "image":
pass
if len(content_text) > 0:
return ActionResponse(
action=Action.REQLLM, result=content_text, response=""
)
except Exception as e:
self.logger.bind(tag=TAG).error(f"MCP工具调用错误: {e}")
return ActionResponse(
action=Action.REQLLM, result="工具调用出错", response=""
)
return ActionResponse(action=Action.REQLLM, result="工具调用出错", response="")
def _handle_function_result(self, result, function_call_data):
if result.action == Action.RESPONSE: # 直接回复前端
text = result.response
@@ -867,7 +786,7 @@ class ConnectionHandler:
)
self.chat(text, tool_call=True)
elif result.action == Action.NOTFOUND or result.action == Action.ERROR:
text = result.result
text = result.response if result.response else result.result
self.tts.tts_one_sentence(self, ContentType.TEXT, content_detail=text)
self.dialogue.put(Message(role="assistant", content=text))
else:
@@ -922,9 +841,9 @@ class ConnectionHandler:
self.timeout_task.cancel()
self.timeout_task = None
# 清理MCP资源
if hasattr(self, "mcp_manager") and self.mcp_manager:
await self.mcp_manager.cleanup_all()
# 清理工具处理器资源
if hasattr(self, "func_handler") and self.func_handler:
await self.func_handler.cleanup()
# 触发停止事件
if self.stop_event:
@@ -976,7 +895,6 @@ class ConnectionHandler:
def reset_vad_states(self):
self.client_audio_buffer = bytearray()
self.client_have_voice = False
self.client_have_voice_last_time = 0
self.client_voice_stop = False
self.logger.bind(tag=TAG).debug("VAD states reset.")
@@ -995,10 +913,18 @@ class ConnectionHandler:
"""检查连接超时"""
try:
while not self.stop_event.is_set():
await asyncio.sleep(self.timeout_seconds)
if not self.stop_event.is_set():
self.logger.bind(tag=TAG).info("连接超时,准备关闭")
await self.close(self.websocket)
break
# 检查是否超时(只有在时间戳已初始化的情况下)
if self.last_activity_time > 0.0:
current_time = time.time() * 1000
if (
current_time - self.last_activity_time
> self.timeout_seconds * 1000
):
if not self.stop_event.is_set():
self.logger.bind(tag=TAG).info("连接超时,准备关闭")
await self.close(self.websocket)
break
# 每10秒检查一次,避免过于频繁
await asyncio.sleep(10)
except Exception as e:
self.logger.bind(tag=TAG).error(f"超时检查任务出错: {e}")
@@ -1,93 +0,0 @@
from config.logger import setup_logging
import json
from plugins_func.register import (
FunctionRegistry,
ActionResponse,
Action,
ToolType,
DeviceTypeRegistry,
)
from plugins_func.functions.hass_init import append_devices_to_prompt
TAG = __name__
class FunctionHandler:
def __init__(self, conn):
self.conn = conn
self.config = conn.config
self.device_type_registry = DeviceTypeRegistry()
self.function_registry = FunctionRegistry()
self.register_nessary_functions()
self.register_config_functions()
self.functions_desc = self.function_registry.get_all_function_desc()
self.finish_init = True
def upload_functions_desc(self):
self.functions_desc = self.function_registry.get_all_function_desc()
def current_support_functions(self):
func_names = []
for func in self.functions_desc:
func_names.append(func["function"]["name"])
# 打印当前支持的函数列表
self.conn.logger.bind(tag=TAG, session_id=self.conn.session_id).info(
f"当前支持的函数列表: {func_names}"
)
return func_names
def get_functions(self):
"""获取功能调用配置"""
return self.functions_desc
def register_nessary_functions(self):
"""注册必要的函数"""
self.function_registry.register_function("handle_exit_intent")
self.function_registry.register_function("get_time")
self.function_registry.register_function("get_lunar")
def register_config_functions(self):
"""注册配置中的函数,可以不同客户端使用不同的配置"""
for func in self.config["Intent"][self.config["selected_module"]["Intent"]].get(
"functions", []
):
self.function_registry.register_function(func)
"""home assistant需要初始化提示词"""
append_devices_to_prompt(self.conn)
def get_function(self, name):
return self.function_registry.get_function(name)
def handle_llm_function_call(self, conn, function_call_data):
try:
function_name = function_call_data["name"]
funcItem = self.get_function(function_name)
if not funcItem:
return ActionResponse(
action=Action.NOTFOUND, result="没有找到对应的函数", response=""
)
func = funcItem.func
arguments = function_call_data["arguments"]
arguments = json.loads(arguments) if arguments else {}
self.conn.logger.bind(tag=TAG).debug(
f"调用函数: {function_name}, 参数: {arguments}"
)
if (
funcItem.type == ToolType.SYSTEM_CTL
or funcItem.type == ToolType.IOT_CTL
):
return func(conn, **arguments)
elif funcItem.type == ToolType.WAIT:
return func(**arguments)
elif funcItem.type == ToolType.CHANGE_SYS_PROMPT:
return func(conn, **arguments)
else:
return ActionResponse(
action=Action.NOTFOUND, result="没有找到对应的函数", response=""
)
except Exception as e:
self.conn.logger.bind(tag=TAG).error(f"处理function call错误: {e}")
return None
@@ -7,7 +7,7 @@ from core.utils.util import audio_to_data
from core.handle.sendAudioHandle import sendAudioMessage, send_stt_message
from core.utils.util import remove_punctuation_and_length, opus_datas_to_wav_bytes
from core.providers.tts.dto.dto import ContentType, SentenceType
from core.handle.mcpHandle import (
from core.providers.tools.device_mcp import (
MCPClient,
send_mcp_initialize_message,
send_mcp_tools_list_request,
@@ -6,7 +6,7 @@ from core.handle.helloHandle import checkWakeupWords
from core.utils.util import remove_punctuation_and_length
from core.providers.tts.dto.dto import ContentType
from core.utils.dialogue import Message
from core.handle.mcpHandle import call_mcp_tool
from core.providers.tools.device_mcp import call_mcp_tool
from plugins_func.register import Action, ActionResponse
from loguru import logger
@@ -106,36 +106,18 @@ async def process_intent_result(conn, intent_result, original_text):
def process_function_call():
conn.dialogue.put(Message(role="user", content=original_text))
# 处理Server端MCP工具调用
if conn.mcp_manager.is_mcp_tool(function_name):
result = conn._handle_mcp_tool_call(function_call_data)
elif hasattr(conn, "mcp_client") and conn.mcp_client.has_tool(
function_name
):
# 如果是小智端MCP工具调用
conn.logger.bind(tag=TAG).debug(
f"调用小智端MCP工具: {function_name}, 参数: {function_args}"
)
try:
result = asyncio.run_coroutine_threadsafe(
call_mcp_tool(
conn, conn.mcp_client, function_name, function_args
),
conn.loop,
).result()
conn.logger.bind(tag=TAG).debug(f"MCP工具调用结果: {result}")
result = ActionResponse(
action=Action.REQLLM, result=result, response=""
)
except Exception as e:
conn.logger.bind(tag=TAG).error(f"MCP工具调用失败: {e}")
result = ActionResponse(
action=Action.REQLLM, result="MCP工具调用失败", response=""
)
else:
# 处理系统函数
result = conn.func_handler.handle_llm_function_call(
conn, function_call_data
# 使用统一工具处理器处理所有工具调用
try:
result = asyncio.run_coroutine_threadsafe(
conn.func_handler.handle_llm_function_call(
conn, function_call_data
),
conn.loop,
).result()
except Exception as e:
conn.logger.bind(tag=TAG).error(f"工具调用失败: {e}")
result = ActionResponse(
action=Action.ERROR, result=str(e), response=str(e)
)
if result:
@@ -1,423 +0,0 @@
import json
import asyncio
from plugins_func.register import (
FunctionItem,
register_device_function,
ActionResponse,
Action,
ToolType,
)
TAG = __name__
def wrap_async_function(async_func):
"""包装异步函数为同步函数"""
def wrapper(*args, **kwargs):
try:
# 获取连接对象(第一个参数)
conn = args[0]
if not hasattr(conn, "loop"):
conn.logger.bind(tag=TAG).error("Connection对象没有loop属性")
return ActionResponse(
Action.ERROR,
"Connection对象没有loop属性",
"执行操作时出错: Connection对象没有loop属性",
)
# 使用conn对象中的事件循环
loop = conn.loop
# 在conn的事件循环中运行异步函数
future = asyncio.run_coroutine_threadsafe(async_func(*args, **kwargs), loop)
# 等待结果返回
return future.result()
except Exception as e:
conn.logger.bind(tag=TAG).error(f"运行异步函数时出错: {e}")
return ActionResponse(Action.ERROR, str(e), f"执行操作时出错: {e}")
return wrapper
def create_iot_function(device_name, method_name, method_info):
"""
根据IOT设备描述生成通用的控制函数
"""
async def iot_control_function(
conn, response_success=None, response_failure=None, **params
):
try:
# 设置默认响应消息
if not response_success:
response_success = "操作成功"
if not response_failure:
response_failure = "操作失败"
# 打印响应参数
conn.logger.bind(tag=TAG).debug(
f"控制函数接收到的响应参数: success='{response_success}', failure='{response_failure}'"
)
# 发送控制命令
await send_iot_conn(conn, device_name, method_name, params)
# 等待一小段时间让状态更新
await asyncio.sleep(0.1)
# 生成结果信息
result = f"{device_name}{method_name}操作执行成功"
# 处理响应中可能的占位符
response = response_success
# 替换{value}占位符
for param_name, param_value in params.items():
# 先尝试直接替换参数值
if "{" + param_name + "}" in response:
response = response.replace(
"{" + param_name + "}", str(param_value)
)
# 如果有{value}占位符,用相关参数替换
if "{value}" in response:
response = response.replace("{value}", str(param_value))
break
return ActionResponse(Action.RESPONSE, result, response)
except Exception as e:
conn.logger.bind(tag=TAG).error(
f"执行{device_name}{method_name}操作失败: {e}"
)
# 操作失败时使用大模型提供的失败响应
response = response_failure
return ActionResponse(Action.ERROR, str(e), response)
return wrap_async_function(iot_control_function)
def create_iot_query_function(device_name, prop_name, prop_info):
"""
根据IOT设备属性创建查询函数
"""
async def iot_query_function(conn, response_success=None, response_failure=None):
try:
# 打印响应参数
conn.logger.bind(tag=TAG).info(
f"查询函数接收到的响应参数: success='{response_success}', failure='{response_failure}'"
)
value = await get_iot_status(conn, device_name, prop_name)
# 查询成功,生成结果
if value is not None:
# 使用大模型提供的成功响应,并替换其中的占位符
response = response_success.replace("{value}", str(value))
return ActionResponse(Action.RESPONSE, str(value), response)
else:
# 查询失败,使用大模型提供的失败响应
response = response_failure
return ActionResponse(Action.ERROR, f"属性{prop_name}不存在", response)
except Exception as e:
conn.logger.bind(tag=TAG).error(
f"查询{device_name}{prop_name}时出错: {e}"
)
# 查询出错时使用大模型提供的失败响应
response = response_failure
return ActionResponse(Action.ERROR, str(e), response)
return wrap_async_function(iot_query_function)
class IotDescriptor:
"""
A class to represent an IoT descriptor.
"""
def __init__(self, name, description, properties, methods):
self.name = name
self.description = description
self.properties = []
self.methods = []
# 根据描述创建属性
if properties is not None:
for key, value in properties.items():
property_item = {}
property_item["name"] = key
property_item["description"] = value["description"]
if value["type"] == "number":
property_item["value"] = 0
elif value["type"] == "boolean":
property_item["value"] = False
else:
property_item["value"] = ""
self.properties.append(property_item)
# 根据描述创建方法
if methods is not None:
for key, value in methods.items():
method = {}
method["description"] = value["description"]
method["name"] = key
# 检查方法是否有参数
if "parameters" in value:
method["parameters"] = {}
for k, v in value["parameters"].items():
method["parameters"][k] = {
"description": v["description"],
"type": v["type"],
}
self.methods.append(method)
def register_device_type(descriptor, device_type_registry):
"""注册设备类型及其功能"""
device_name = descriptor["name"]
type_id = device_type_registry.generate_device_type_id(descriptor)
# 如果该类型已注册,直接返回类型ID
if type_id in device_type_registry.type_functions:
return type_id
functions = {}
# 为每个属性创建查询函数
for prop_name, prop_info in descriptor["properties"].items():
func_name = f"get_{device_name.lower()}_{prop_name.lower()}"
func_desc = {
"type": "function",
"function": {
"name": func_name,
"description": f"查询{descriptor['description']}{prop_info['description']}",
"parameters": {
"type": "object",
"properties": {
"response_success": {
"type": "string",
"description": f"查询成功时的友好回复,必须使用{{value}}作为占位符表示查询到的值",
},
"response_failure": {
"type": "string",
"description": f"查询失败时的友好回复,例如:'无法获取{device_name}{prop_info['description']}'",
},
},
"required": ["response_success", "response_failure"],
},
},
}
query_func = create_iot_query_function(device_name, prop_name, prop_info)
decorated_func = register_device_function(
func_name, func_desc, ToolType.IOT_CTL
)(query_func)
functions[func_name] = FunctionItem(
func_name, func_desc, decorated_func, ToolType.IOT_CTL
)
# 为每个方法创建控制函数
for method_name, method_info in descriptor["methods"].items():
func_name = f"{device_name.lower()}_{method_name.lower()}"
# 创建参数字典,添加原有参数
parameters = {}
required_params = []
# 如果方法有参数,则添加参数信息
if "parameters" in method_info:
parameters = {
param_name: {
"type": param_info["type"],
"description": param_info["description"],
}
for param_name, param_info in method_info["parameters"].items()
}
required_params = list(method_info["parameters"].keys())
# 添加响应参数
parameters.update(
{
"response_success": {
"type": "string",
"description": "操作成功时的友好回复,关于该设备的操作结果,设备名称尽量使用description中的名称",
},
"response_failure": {
"type": "string",
"description": "操作失败时的友好回复,关于该设备的操作结果,设备名称尽量使用description中的名称",
},
}
)
# 构建必须参数列表(原有参数 + 响应参数)
required_params.extend(["response_success", "response_failure"])
func_desc = {
"type": "function",
"function": {
"name": func_name,
"description": f"{descriptor['description']} - {method_info['description']}",
"parameters": {
"type": "object",
"properties": parameters,
"required": required_params,
},
},
}
control_func = create_iot_function(device_name, method_name, method_info)
decorated_func = register_device_function(
func_name, func_desc, ToolType.IOT_CTL
)(control_func)
functions[func_name] = FunctionItem(
func_name, func_desc, decorated_func, ToolType.IOT_CTL
)
device_type_registry.register_device_type(type_id, functions)
return type_id
# 用于接受前端设备推送的搜索iot描述
async def handleIotDescriptors(conn, descriptors):
wait_max_time = 5
while conn.func_handler is None or not conn.func_handler.finish_init:
await asyncio.sleep(1)
wait_max_time -= 1
if wait_max_time <= 0:
conn.logger.bind(tag=TAG).debug("连接对象没有func_handler")
return
"""处理物联网描述"""
functions_changed = False
for descriptor in descriptors:
# 如果descriptor没有properties和methods,则直接跳过
if "properties" not in descriptor and "methods" not in descriptor:
continue
# 处理缺失properties的情况
if "properties" not in descriptor:
descriptor["properties"] = {}
# 从methods中提取所有参数作为properties
if "methods" in descriptor:
for method_name, method_info in descriptor["methods"].items():
if "parameters" in method_info:
for param_name, param_info in method_info["parameters"].items():
# 将参数信息转换为属性信息
descriptor["properties"][param_name] = {
"description": param_info["description"],
"type": param_info["type"],
}
# 创建IOT设备描述符
iot_descriptor = IotDescriptor(
descriptor["name"],
descriptor["description"],
descriptor["properties"],
descriptor["methods"],
)
conn.iot_descriptors[descriptor["name"]] = iot_descriptor
if conn.load_function_plugin:
# 注册或获取设备类型
device_type_registry = conn.func_handler.device_type_registry
type_id = register_device_type(descriptor, device_type_registry)
device_functions = device_type_registry.get_device_functions(type_id)
# 在连接级注册设备函数
if hasattr(conn, "func_handler"):
for func_name, func_item in device_functions.items():
conn.func_handler.function_registry.register_function(
func_name, func_item
)
conn.logger.bind(tag=TAG).info(
f"注册IOT函数到function handler: {func_name}"
)
functions_changed = True
# 如果注册了新函数,更新function描述列表
if functions_changed and hasattr(conn, "func_handler"):
conn.func_handler.upload_functions_desc()
func_names = conn.func_handler.current_support_functions()
conn.logger.bind(tag=TAG).info(f"设备类型: {type_id}")
conn.logger.bind(tag=TAG).info(
f"更新function描述列表完成,当前支持的函数: {func_names}"
)
async def handleIotStatus(conn, states):
"""处理物联网状态"""
for state in states:
for key, value in conn.iot_descriptors.items():
if key == state["name"]:
for property_item in value.properties:
for k, v in state["state"].items():
if property_item["name"] == k:
if type(v) != type(property_item["value"]):
conn.logger.bind(tag=TAG).error(
f"属性{property_item['name']}的值类型不匹配"
)
break
else:
property_item["value"] = v
conn.logger.bind(tag=TAG).info(
f"物联网状态更新: {key} , {property_item['name']} = {v}"
)
break
break
async def get_iot_status(conn, name, property_name):
"""获取物联网状态"""
for key, value in conn.iot_descriptors.items():
if key == name:
for property_item in value.properties:
if property_item["name"] == property_name:
return property_item["value"]
conn.logger.bind(tag=TAG).warning(f"未找到设备 {name} 的属性 {property_name}")
return None
async def set_iot_status(conn, name, property_name, value):
"""设置物联网状态"""
for key, iot_descriptor in conn.iot_descriptors.items():
if key == name:
for property_item in iot_descriptor.properties:
if property_item["name"] == property_name:
if type(value) != type(property_item["value"]):
conn.logger.bind(tag=TAG).error(
f"属性{property_item['name']}的值类型不匹配"
)
return
property_item["value"] = value
conn.logger.bind(tag=TAG).info(
f"物联网状态更新: {name} , {property_name} = {value}"
)
return
conn.logger.bind(tag=TAG).warning(f"未找到设备 {name} 的属性 {property_name}")
async def send_iot_conn(conn, name, method_name, parameters):
"""发送物联网指令"""
for key, value in conn.iot_descriptors.items():
if key == name:
# 找到了设备
for method in value.methods:
# 找到了方法
if method["name"] == method_name:
# 构建命令对象
command = {
"name": name,
"method": method_name,
}
# 只有当参数不为空时才添加parameters字段
if parameters:
command["parameters"] = parameters
send_message = json.dumps({"type": "iot", "commands": [command]})
await conn.websocket.send(send_message)
conn.logger.bind(tag=TAG).info(f"发送物联网指令: {send_message}")
return
conn.logger.bind(tag=TAG).error(f"未找到方法{method_name}")
@@ -66,12 +66,11 @@ async def startToChat(conn, text):
async def no_voice_close_connect(conn, have_voice):
if have_voice:
conn.client_no_voice_last_time = 0.0
conn.last_activity_time = time.time() * 1000
return
if conn.client_no_voice_last_time == 0.0:
conn.client_no_voice_last_time = time.time() * 1000
else:
no_voice_time = time.time() * 1000 - conn.client_no_voice_last_time
# 只有在已经初始化过时间戳的情况下才进行超时检查
if conn.last_activity_time > 0.0:
no_voice_time = time.time() * 1000 - conn.last_activity_time
close_connection_no_voice_time = int(
conn.config.get("close_connection_no_voice_time", 120)
)
@@ -92,10 +92,8 @@ async def sendAudio(conn, audios, pre_buffer=True):
if conn.client_abort:
break
# 每分钟重置一次计时器
if time.perf_counter() - last_reset_time > 60:
await conn.reset_timeout()
last_reset_time = time.perf_counter()
# 重置没有声音的状态
conn.last_activity_time = time.time() * 1000
# 计算预期发送时间
expected_time = start_time + (play_position / 1000)
@@ -1,11 +1,11 @@
import json
from core.handle.abortHandle import handleAbortMessage
from core.handle.helloHandle import handleHelloMessage
from core.handle.mcpHandle import handle_mcp_message
from core.providers.tools.device_mcp import handle_mcp_message
from core.utils.util import remove_punctuation_and_length, filter_sensitive_info
from core.handle.receiveAudioHandle import startToChat, handleAudioMessage
from core.handle.sendAudioHandle import send_stt_message, send_tts_message
from core.handle.iotHandle import handleIotDescriptors, handleIotStatus
from core.providers.tools.device_iot import handleIotDescriptors, handleIotStatus
from core.handle.reportHandle import enqueue_asr_report
import asyncio
@@ -77,7 +77,7 @@ async def handleTextMessage(conn, message):
if "states" in msg_json:
asyncio.create_task(handleIotStatus(conn, msg_json["states"]))
elif msg_json["type"] == "mcp":
conn.logger.bind(tag=TAG).info(f"收到mcp消息:{message}")
conn.logger.bind(tag=TAG).info(f"收到mcp消息:{message[:100]}")
if "payload" in msg_json:
asyncio.create_task(
handle_mcp_message(conn, conn.mcp_client, msg_json["payload"])
-131
View File
@@ -1,131 +0,0 @@
"""MCP服务管理器"""
import asyncio
import os, json
from typing import Dict, Any, List
from .MCPClient import MCPClient
from plugins_func.register import register_function, ToolType
from config.config_loader import get_project_dir
TAG = __name__
class MCPManager:
"""管理多个MCP服务的集中管理器"""
def __init__(self, conn) -> None:
"""
初始化MCP管理器
"""
self.conn = conn
self.config_path = get_project_dir() + "data/.mcp_server_settings.json"
if os.path.exists(self.config_path) == False:
self.config_path = ""
self.conn.logger.bind(tag=TAG).warning(
f"请检查mcp服务配置文件:data/.mcp_server_settings.json"
)
self.client: Dict[str, MCPClient] = {}
self.tools = []
def load_config(self) -> Dict[str, Any]:
"""加载MCP服务配置
Returns:
Dict[str, Any]: 服务配置字典
"""
if len(self.config_path) == 0:
return {}
try:
with open(self.config_path, "r", encoding="utf-8") as f:
config = json.load(f)
return config.get("mcpServers", {})
except Exception as e:
self.conn.logger.bind(tag=TAG).error(
f"Error loading MCP config from {self.config_path}: {e}"
)
return {}
async def initialize_servers(self) -> None:
"""初始化所有MCP服务"""
config = self.load_config()
for name, srv_config in config.items():
if not srv_config.get("command") and not srv_config.get("url"):
self.conn.logger.bind(tag=TAG).warning(
f"Skipping server {name}: neither command nor url specified"
)
continue
try:
client = MCPClient(srv_config)
await client.initialize()
self.client[name] = client
self.conn.logger.bind(tag=TAG).info(f"Initialized MCP client: {name}")
client_tools = client.get_available_tools()
self.tools.extend(client_tools)
for tool in client_tools:
func_name = "mcp_" + tool["function"]["name"]
register_function(func_name, tool, ToolType.MCP_CLIENT)(
self.execute_tool
)
self.conn.func_handler.function_registry.register_function(
func_name
)
except Exception as e:
self.conn.logger.bind(tag=TAG).error(
f"Failed to initialize MCP server {name}: {e}"
)
self.conn.func_handler.upload_functions_desc()
def get_all_tools(self) -> List[Dict[str, Any]]:
"""获取所有服务的工具function定义
Returns:
List[Dict[str, Any]]: 所有工具的function定义列表
"""
return self.tools
def is_mcp_tool(self, tool_name: str) -> bool:
"""检查是否是MCP工具
Args:
tool_name: 工具名称
Returns:
bool: 是否是MCP工具
"""
for tool in self.tools:
if (
tool.get("function") != None
and tool["function"].get("name") == tool_name
):
return True
return False
async def execute_tool(self, tool_name: str, arguments: Dict[str, Any]) -> Any:
"""执行工具调用
Args:
tool_name: 工具名称
arguments: 工具参数
Returns:
Any: 工具执行结果
Raises:
ValueError: 工具未找到时抛出
"""
self.conn.logger.bind(tag=TAG).info(
f"Executing tool {tool_name} with arguments: {arguments}"
)
for client in self.client.values():
if client.has_tool(tool_name):
return await client.call_tool(tool_name, arguments)
raise ValueError(f"Tool {tool_name} not found in any MCP server")
async def cleanup_all(self) -> None:
"""依次关闭所有 MCPClient,不让异常阻断整体流程。"""
for name, client in list(self.client.items()):
try:
await asyncio.wait_for(client.cleanup(), timeout=20)
self.conn.logger.bind(tag=TAG).info(f"MCP client closed: {name}")
except (asyncio.TimeoutError, Exception) as e:
self.conn.logger.bind(tag=TAG).error(
f"Error closing MCP client {name}: {e}"
)
self.client.clear()
@@ -93,9 +93,7 @@ class ASRProvider(ASRProviderBase):
# 检查初始化响应
if "code" in result and result["code"] != 1000:
error_msg = f"ASR服务初始化失败: {result.get('payload_msg', {}).get('message', '未知错误')}"
if "payload_msg" in result:
error_msg += f"\n详细错误信息: {json.dumps(result['payload_msg'], ensure_ascii=False)}"
error_msg = f"ASR服务初始化失败: {result.get('payload_msg', {}).get('error', '未知错误')}"
logger.bind(tag=TAG).error(error_msg)
raise Exception(error_msg)
@@ -256,7 +254,6 @@ class ASRProvider(ASRProviderBase):
"X-Api-Access-Key": self.access_token,
"X-Api-Resource-Id": "volc.bigasr.sauc.duration",
"X-Api-Connect-Id": str(uuid.uuid4()),
"Host": "openspeech.bytedance.com",
}
def generate_header(
@@ -309,9 +306,14 @@ class ASRProvider(ASRProviderBase):
# 如果是错误响应
if message_type == 0x0F: # SERVER_ERROR_RESPONSE
code = int.from_bytes(header[4:8], "big", signed=False)
error_msg = res[8:].decode("utf-8")
return {"code": code, "error": error_msg}
code = int.from_bytes(res[4:8], "big", signed=False)
msg_length = int.from_bytes(res[8:12], "big", signed=False)
error_msg = json.loads(res[12:].decode("utf-8"))
return {
"code": code,
"msg_length": msg_length,
"payload_msg": error_msg,
}
# 获取JSON数据(跳过12字节头部)
try:
@@ -95,6 +95,10 @@ class IntentProvider(IntentProviderBase):
"1. 只返回JSON格式,不要包含任何其他文字\n"
'2. 如果没有找到匹配的函数,返回{"function_call": {"name": "continue_chat"}}\n'
"3. 确保返回的JSON格式正确,包含所有必要的字段\n"
"特殊说明:\n"
"- 当用户单次输入包含多个指令时(如'打开灯并且调高音量'\n"
"- 请返回多个function_call组成的JSON数组\n"
"- 示例:{'function_calls': [{name:'light_on'}, {name:'volume_up'}]}"
)
return prompt
@@ -1,3 +1,4 @@
import httpx
import openai
from openai.types import CompletionUsage
from config.logger import setup_logging
@@ -16,6 +17,9 @@ class LLMProvider(LLMProviderBase):
self.base_url = config.get("base_url")
else:
self.base_url = config.get("url")
# 增加timeout的配置项,单位为秒
timeout = config.get("timeout", 300)
self.timeout = int(timeout) if timeout else 300
param_defaults = {
"max_tokens": (500, int),
@@ -42,7 +46,7 @@ class LLMProvider(LLMProviderBase):
model_key_msg = check_model_key("LLM", self.api_key)
if model_key_msg:
logger.bind(tag=TAG).error(model_key_msg)
self.client = openai.OpenAI(api_key=self.api_key, base_url=self.base_url)
self.client = openai.OpenAI(api_key=self.api_key, base_url=self.base_url, timeout=httpx.Timeout(self.timeout))
def response(self, session_id, dialogue, **kwargs):
try:
@@ -5,6 +5,7 @@ import os
import yaml
from config.config_loader import get_project_dir
from config.manage_api_client import save_mem_local_short
from core.utils.util import check_model_key
short_term_memory_prompt = """
@@ -145,6 +146,10 @@ class MemoryProvider(MemoryProviderBase):
# 打印使用的模型信息
model_info = getattr(self.llm, "model_name", str(self.llm.__class__.__name__))
logger.bind(tag=TAG).debug(f"使用记忆保存模型: {model_info}")
api_key = getattr(self.llm, "api_key", None)
memory_key_msg = check_model_key("记忆总结专用LLM", api_key)
if memory_key_msg:
logger.bind(tag=TAG).error(memory_key_msg)
if self.llm is None:
logger.bind(tag=TAG).error("LLM is not set for memory provider")
return None
@@ -0,0 +1 @@
@@ -0,0 +1,6 @@
"""基础工具定义模块"""
from .tool_types import ToolType, ToolDefinition
from .tool_executor import ToolExecutor
__all__ = ["ToolType", "ToolDefinition", "ToolExecutor"]
@@ -0,0 +1,27 @@
"""工具执行器基类定义"""
from abc import ABC, abstractmethod
from typing import Dict, Any
from .tool_types import ToolDefinition
from plugins_func.register import ActionResponse
class ToolExecutor(ABC):
"""工具执行器抽象基类"""
@abstractmethod
async def execute(
self, conn, tool_name: str, arguments: Dict[str, Any]
) -> ActionResponse:
"""执行工具调用"""
pass
@abstractmethod
def get_tools(self) -> Dict[str, ToolDefinition]:
"""获取该执行器管理的所有工具"""
pass
@abstractmethod
def has_tool(self, tool_name: str) -> bool:
"""检查是否有指定工具"""
pass
@@ -0,0 +1,27 @@
"""工具系统的类型定义"""
from enum import Enum
from dataclasses import dataclass
from typing import Any, Dict, Optional
from plugins_func.register import Action
class ToolType(Enum):
"""工具类型枚举"""
SERVER_PLUGIN = "server_plugin" # 服务端插件
SERVER_MCP = "server_mcp" # 服务端MCP
DEVICE_IOT = "device_iot" # 设备端IoT
DEVICE_MCP = "device_mcp" # 设备端MCP
MCP_ENDPOINT = "mcp_endpoint" # MCP接入点
@dataclass
class ToolDefinition:
"""工具定义"""
name: str # 工具名称
description: Dict[str, Any] # 工具描述(OpenAI函数调用格式)
tool_type: ToolType # 工具类型
parameters: Optional[Dict[str, Any]] = None # 额外参数
@@ -0,0 +1,12 @@
"""设备端IoT工具模块"""
from .iot_descriptor import IotDescriptor
from .iot_handler import handleIotDescriptors, handleIotStatus
from .iot_executor import DeviceIoTExecutor
__all__ = [
"IotDescriptor",
"handleIotDescriptors",
"handleIotStatus",
"DeviceIoTExecutor",
]
@@ -0,0 +1,46 @@
"""IoT设备描述符定义"""
from config.logger import setup_logging
TAG = __name__
logger = setup_logging()
class IotDescriptor:
"""IoT设备描述符"""
def __init__(self, name, description, properties, methods):
self.name = name
self.description = description
self.properties = []
self.methods = []
# 根据描述创建属性
if properties is not None:
for key, value in properties.items():
property_item = {}
property_item["name"] = key
property_item["description"] = value["description"]
if value["type"] == "number":
property_item["value"] = 0
elif value["type"] == "boolean":
property_item["value"] = False
else:
property_item["value"] = ""
self.properties.append(property_item)
# 根据描述创建方法
if methods is not None:
for key, value in methods.items():
method = {}
method["description"] = value["description"]
method["name"] = key
# 检查方法是否有参数
if "parameters" in value:
method["parameters"] = {}
for k, v in value["parameters"].items():
method["parameters"][k] = {
"description": v["description"],
"type": v["type"],
}
self.methods.append(method)
@@ -0,0 +1,238 @@
"""设备端IoT工具执行器"""
import json
import asyncio
from typing import Dict, Any
from ..base import ToolType, ToolDefinition, ToolExecutor
from plugins_func.register import Action, ActionResponse
class DeviceIoTExecutor(ToolExecutor):
"""设备端IoT工具执行器"""
def __init__(self, conn):
self.conn = conn
self.iot_tools: Dict[str, ToolDefinition] = {}
async def execute(
self, conn, tool_name: str, arguments: Dict[str, Any]
) -> ActionResponse:
"""执行设备端IoT工具"""
if not self.has_tool(tool_name):
return ActionResponse(
action=Action.NOTFOUND, response=f"IoT工具 {tool_name} 不存在"
)
try:
# 解析工具名称,获取设备名和操作类型
if tool_name.startswith("get_"):
# 查询操作:get_devicename_property
parts = tool_name.split("_", 2)
if len(parts) >= 3:
device_name = parts[1]
property_name = parts[2]
value = await self._get_iot_status(device_name, property_name)
if value is not None:
# 处理响应模板
response_success = arguments.get(
"response_success", "查询成功:{value}"
)
response = response_success.replace("{value}", str(value))
return ActionResponse(
action=Action.RESPONSE,
response=response,
)
else:
response_failure = arguments.get(
"response_failure", f"无法获取{device_name}的状态"
)
return ActionResponse(
action=Action.ERROR, response=response_failure
)
else:
# 控制操作:devicename_method
parts = tool_name.split("_", 1)
if len(parts) >= 2:
device_name = parts[0]
method_name = parts[1]
# 提取控制参数(排除响应参数)
control_params = {
k: v
for k, v in arguments.items()
if k not in ["response_success", "response_failure"]
}
# 发送IoT控制命令
await self._send_iot_command(
device_name, method_name, control_params
)
# 等待状态更新
await asyncio.sleep(0.1)
response_success = arguments.get("response_success", "操作成功")
# 处理响应中的占位符
for param_name, param_value in control_params.items():
placeholder = "{" + param_name + "}"
if placeholder in response_success:
response_success = response_success.replace(
placeholder, str(param_value)
)
if "{value}" in response_success:
response_success = response_success.replace(
"{value}", str(param_value)
)
break
return ActionResponse(
action=Action.REQLLM,
result=response_success,
)
return ActionResponse(action=Action.ERROR, response="无法解析IoT工具名称")
except Exception as e:
response_failure = arguments.get("response_failure", "操作失败")
return ActionResponse(action=Action.ERROR, response=response_failure)
async def _get_iot_status(self, device_name: str, property_name: str):
"""获取IoT设备状态"""
for key, value in self.conn.iot_descriptors.items():
if key.lower() == device_name.lower():
for property_item in value.properties:
if property_item["name"].lower() == property_name.lower():
return property_item["value"]
return None
async def _send_iot_command(
self, device_name: str, method_name: str, parameters: Dict[str, Any]
):
"""发送IoT控制命令"""
for key, value in self.conn.iot_descriptors.items():
if key.lower() == device_name.lower():
for method in value.methods:
if method["name"].lower() == method_name.lower():
command = {
"name": key,
"method": method["name"],
}
if parameters:
command["parameters"] = parameters
send_message = json.dumps(
{"type": "iot", "commands": [command]}
)
await self.conn.websocket.send(send_message)
return
raise Exception(f"未找到设备{device_name}的方法{method_name}")
def register_iot_tools(self, descriptors: list):
"""注册IoT工具"""
for descriptor in descriptors:
device_name = descriptor["name"]
device_desc = descriptor["description"]
# 注册查询工具
if "properties" in descriptor:
for prop_name, prop_info in descriptor["properties"].items():
tool_name = f"get_{device_name.lower()}_{prop_name.lower()}"
tool_desc = {
"type": "function",
"function": {
"name": tool_name,
"description": f"查询{device_desc}{prop_info['description']}",
"parameters": {
"type": "object",
"properties": {
"response_success": {
"type": "string",
"description": f"查询成功时的友好回复,必须使用{{value}}作为占位符表示查询到的值",
},
"response_failure": {
"type": "string",
"description": f"查询失败时的友好回复",
},
},
"required": ["response_success", "response_failure"],
},
},
}
self.iot_tools[tool_name] = ToolDefinition(
name=tool_name,
description=tool_desc,
tool_type=ToolType.DEVICE_IOT,
)
# 注册控制工具
if "methods" in descriptor:
for method_name, method_info in descriptor["methods"].items():
tool_name = f"{device_name.lower()}_{method_name.lower()}"
# 构建参数
parameters = {}
required_params = []
# 添加方法的原始参数
if "parameters" in method_info:
parameters.update(
{
param_name: {
"type": param_info["type"],
"description": param_info["description"],
}
for param_name, param_info in method_info[
"parameters"
].items()
}
)
required_params.extend(method_info["parameters"].keys())
# 添加响应参数
parameters.update(
{
"response_success": {
"type": "string",
"description": "操作成功时的友好回复",
},
"response_failure": {
"type": "string",
"description": "操作失败时的友好回复",
},
}
)
required_params.extend(["response_success", "response_failure"])
tool_desc = {
"type": "function",
"function": {
"name": tool_name,
"description": f"{device_desc} - {method_info['description']}",
"parameters": {
"type": "object",
"properties": parameters,
"required": required_params,
},
},
}
self.iot_tools[tool_name] = ToolDefinition(
name=tool_name,
description=tool_desc,
tool_type=ToolType.DEVICE_IOT,
)
def get_tools(self) -> Dict[str, ToolDefinition]:
"""获取所有设备端IoT工具"""
return self.iot_tools.copy()
def has_tool(self, tool_name: str) -> bool:
"""检查是否有指定的设备端IoT工具"""
return tool_name in self.iot_tools
@@ -0,0 +1,83 @@
"""IoT设备支持模块,提供IoT设备描述符和状态处理"""
import asyncio
from config.logger import setup_logging
from .iot_descriptor import IotDescriptor
TAG = __name__
logger = setup_logging()
async def handleIotDescriptors(conn, descriptors):
"""处理物联网描述"""
wait_max_time = 5
while (
not hasattr(conn, "func_handler")
or conn.func_handler is None
or not conn.func_handler.finish_init
):
await asyncio.sleep(1)
wait_max_time -= 1
if wait_max_time <= 0:
logger.bind(tag=TAG).debug("连接对象没有func_handler")
return
functions_changed = False
for descriptor in descriptors:
# 如果descriptor没有properties和methods,则直接跳过
if "properties" not in descriptor and "methods" not in descriptor:
continue
# 处理缺失properties的情况
if "properties" not in descriptor:
descriptor["properties"] = {}
# 从methods中提取所有参数作为properties
if "methods" in descriptor:
for method_name, method_info in descriptor["methods"].items():
if "parameters" in method_info:
for param_name, param_info in method_info["parameters"].items():
# 将参数信息转换为属性信息
descriptor["properties"][param_name] = {
"description": param_info["description"],
"type": param_info["type"],
}
# 创建IOT设备描述符
iot_descriptor = IotDescriptor(
descriptor["name"],
descriptor["description"],
descriptor["properties"],
descriptor["methods"],
)
conn.iot_descriptors[descriptor["name"]] = iot_descriptor
functions_changed = True
# 如果注册了新函数,更新function描述列表
if functions_changed and hasattr(conn, "func_handler"):
# 注册IoT工具到统一工具处理器
await conn.func_handler.register_iot_tools(descriptors)
conn.func_handler.current_support_functions()
async def handleIotStatus(conn, states):
"""处理物联网状态"""
for state in states:
for key, value in conn.iot_descriptors.items():
if key == state["name"]:
for property_item in value.properties:
for k, v in state["state"].items():
if property_item["name"] == k:
if type(v) != type(property_item["value"]):
logger.bind(tag=TAG).error(
f"属性{property_item['name']}的值类型不匹配"
)
break
else:
property_item["value"] = v
logger.bind(tag=TAG).info(
f"物联网状态更新: {key} , {property_item['name']} = {v}"
)
break
break
@@ -0,0 +1,21 @@
"""设备端MCP工具模块"""
from .mcp_client import MCPClient
from .mcp_handler import (
send_mcp_message,
handle_mcp_message,
send_mcp_initialize_message,
send_mcp_tools_list_request,
call_mcp_tool,
)
from .mcp_executor import DeviceMCPExecutor
__all__ = [
"MCPClient",
"send_mcp_message",
"handle_mcp_message",
"send_mcp_initialize_message",
"send_mcp_tools_list_request",
"call_mcp_tool",
"DeviceMCPExecutor",
]
@@ -0,0 +1,93 @@
"""设备端MCP客户端定义"""
import asyncio
from concurrent.futures import Future
from core.utils.util import sanitize_tool_name
from config.logger import setup_logging
TAG = __name__
logger = setup_logging()
class MCPClient:
"""设备端MCP客户端,用于管理MCP状态和工具"""
def __init__(self):
self.tools = {} # sanitized_name -> tool_data
self.name_mapping = {}
self.ready = False
self.call_results = {} # To store Futures for tool call responses
self.next_id = 1
self.lock = asyncio.Lock()
self._cached_available_tools = None # Cache for get_available_tools
def has_tool(self, name: str) -> bool:
return name in self.tools
def get_available_tools(self) -> list:
# Check if the cache is valid
if self._cached_available_tools is not None:
return self._cached_available_tools
# If cache is not valid, regenerate the list
result = []
for tool_name, tool_data in self.tools.items():
function_def = {
"name": tool_name,
"description": tool_data["description"],
"parameters": {
"type": tool_data["inputSchema"].get("type", "object"),
"properties": tool_data["inputSchema"].get("properties", {}),
"required": tool_data["inputSchema"].get("required", []),
},
}
result.append({"type": "function", "function": function_def})
self._cached_available_tools = result # Store the generated list in cache
return result
async def is_ready(self) -> bool:
async with self.lock:
return self.ready
async def set_ready(self, status: bool):
async with self.lock:
self.ready = status
async def add_tool(self, tool_data: dict):
async with self.lock:
sanitized_name = sanitize_tool_name(tool_data["name"])
self.tools[sanitized_name] = tool_data
self.name_mapping[sanitized_name] = tool_data["name"]
self._cached_available_tools = (
None # Invalidate the cache when a tool is added
)
async def get_next_id(self) -> int:
async with self.lock:
current_id = self.next_id
self.next_id += 1
return current_id
async def register_call_result_future(self, id: int, future: Future):
async with self.lock:
self.call_results[id] = future
async def resolve_call_result(self, id: int, result: any):
async with self.lock:
if id in self.call_results:
future = self.call_results.pop(id)
if not future.done():
future.set_result(result)
async def reject_call_result(self, id: int, exception: Exception):
async with self.lock:
if id in self.call_results:
future = self.call_results.pop(id)
if not future.done():
future.set_exception(exception)
async def cleanup_call_result(self, id: int):
async with self.lock:
if id in self.call_results:
self.call_results.pop(id)
@@ -0,0 +1,89 @@
"""设备端MCP工具执行器"""
from typing import Dict, Any
from ..base import ToolType, ToolDefinition, ToolExecutor
from plugins_func.register import Action, ActionResponse
from .mcp_handler import call_mcp_tool
class DeviceMCPExecutor(ToolExecutor):
"""设备端MCP工具执行器"""
def __init__(self, conn):
self.conn = conn
async def execute(
self, conn, tool_name: str, arguments: Dict[str, Any]
) -> ActionResponse:
"""执行设备端MCP工具"""
if not hasattr(conn, "mcp_client") or not conn.mcp_client:
return ActionResponse(
action=Action.ERROR,
response="设备端MCP客户端未初始化",
)
if not await conn.mcp_client.is_ready():
return ActionResponse(
action=Action.ERROR,
response="设备端MCP客户端未准备就绪",
)
try:
# 转换参数为JSON字符串
import json
args_str = json.dumps(arguments) if arguments else "{}"
# 调用设备端MCP工具
result = await call_mcp_tool(conn, conn.mcp_client, tool_name, args_str)
resultJson = None
if isinstance(result, str):
try:
resultJson = json.loads(result)
except Exception as e:
pass
# 视觉大模型不经过二次LLM处理
if (
resultJson is not None
and isinstance(resultJson, dict)
and "action" in resultJson
):
return ActionResponse(
action=Action[resultJson["action"]],
response=resultJson.get("response", ""),
)
return ActionResponse(action=Action.REQLLM, result=str(result))
except ValueError as e:
return ActionResponse(action=Action.NOTFOUND, response=str(e))
except Exception as e:
return ActionResponse(action=Action.ERROR, response=str(e))
def get_tools(self) -> Dict[str, ToolDefinition]:
"""获取所有设备端MCP工具"""
if not hasattr(self.conn, "mcp_client") or not self.conn.mcp_client:
return {}
tools = {}
mcp_tools = self.conn.mcp_client.get_available_tools()
for tool in mcp_tools:
func_def = tool.get("function", {})
tool_name = func_def.get("name", "")
if tool_name:
tools[tool_name] = ToolDefinition(
name=tool_name, description=tool, tool_type=ToolType.DEVICE_MCP
)
return tools
def has_tool(self, tool_name: str) -> bool:
"""检查是否有指定的设备端MCP工具"""
if not hasattr(self.conn, "mcp_client") or not self.conn.mcp_client:
return False
return self.conn.mcp_client.has_tool(tool_name)
@@ -1,14 +1,19 @@
"""设备端MCP客户端支持模块"""
import json
import asyncio
import re
from concurrent.futures import Future
from core.utils.util import get_vision_url, sanitize_tool_name
from core.utils.auth import AuthToken
from config.logger import setup_logging
TAG = __name__
logger = setup_logging()
class MCPClient:
"""MCPClient,用于管理MCP状态和工具"""
"""设备端MCP客户端,用于管理MCP状态和工具"""
def __init__(self):
self.tools = {} # sanitized_name -> tool_data
@@ -94,24 +99,24 @@ class MCPClient:
async def send_mcp_message(conn, payload: dict):
"""Helper to send MCP messages, encapsulating common logic."""
if not conn.features.get("mcp"):
conn.logger.bind(tag=TAG).warning("客户端不支持MCP,无法发送MCP消息")
logger.bind(tag=TAG).warning("客户端不支持MCP,无法发送MCP消息")
return
message = json.dumps({"type": "mcp", "payload": payload})
try:
await conn.websocket.send(message)
conn.logger.bind(tag=TAG).info(f"成功发送MCP消息: {message}")
logger.bind(tag=TAG).info(f"成功发送MCP消息: {message}")
except Exception as e:
conn.logger.bind(tag=TAG).error(f"发送MCP消息失败: {e}")
logger.bind(tag=TAG).error(f"发送MCP消息失败: {e}")
async def handle_mcp_message(conn, mcp_client: MCPClient, payload: dict):
"""处理MCP消息,包括初始化、工具列表和工具调用响应等"""
conn.logger.bind(tag=TAG).info(f"处理MCP消息: {payload}")
logger.bind(tag=TAG).info(f"处理MCP消息: {str(payload)[:100]}")
if not isinstance(payload, dict):
conn.logger.bind(tag=TAG).error("MCP消息缺少payload字段或格式错误")
logger.bind(tag=TAG).error("MCP消息缺少payload字段或格式错误")
return
# Handle result
@@ -121,32 +126,32 @@ async def handle_mcp_message(conn, mcp_client: MCPClient, payload: dict):
# Check for tool call response first
if msg_id in mcp_client.call_results:
conn.logger.bind(tag=TAG).debug(
logger.bind(tag=TAG).debug(
f"收到工具调用响应,ID: {msg_id}, 结果: {result}"
)
await mcp_client.resolve_call_result(msg_id, result)
return
if msg_id == 1: # mcpInitializeID
conn.logger.bind(tag=TAG).debug("收到MCP初始化响应")
logger.bind(tag=TAG).debug("收到MCP初始化响应")
server_info = result.get("serverInfo")
if isinstance(server_info, dict):
name = server_info.get("name")
version = server_info.get("version")
conn.logger.bind(tag=TAG).info(
logger.bind(tag=TAG).info(
f"客户端MCP服务器信息: name={name}, version={version}"
)
return
elif msg_id == 2: # mcpToolsListID
conn.logger.bind(tag=TAG).debug("收到MCP工具列表响应")
logger.bind(tag=TAG).debug("收到MCP工具列表响应")
if isinstance(result, dict) and "tools" in result:
tools_data = result["tools"]
if not isinstance(tools_data, list):
conn.logger.bind(tag=TAG).error("工具列表格式错误")
logger.bind(tag=TAG).error("工具列表格式错误")
return
conn.logger.bind(tag=TAG).info(
logger.bind(tag=TAG).info(
f"客户端设备支持的工具数量: {len(tools_data)}"
)
@@ -172,7 +177,7 @@ async def handle_mcp_message(conn, mcp_client: MCPClient, payload: dict):
"inputSchema": input_schema,
}
await mcp_client.add_tool(new_tool)
conn.logger.bind(tag=TAG).debug(f"客户端工具 #{i+1}: {name}")
logger.bind(tag=TAG).debug(f"客户端工具 #{i+1}: {name}")
# 替换所有工具描述中的工具名称
for tool_data in mcp_client.tools.values():
@@ -190,24 +195,27 @@ async def handle_mcp_message(conn, mcp_client: MCPClient, payload: dict):
next_cursor = result.get("nextCursor", "")
if next_cursor:
conn.logger.bind(tag=TAG).info(
f"有更多工具,nextCursor: {next_cursor}"
)
logger.bind(tag=TAG).info(f"有更多工具,nextCursor: {next_cursor}")
await send_mcp_tools_list_continue_request(conn, next_cursor)
else:
await mcp_client.set_ready(True)
conn.logger.bind(tag=TAG).info("所有工具已获取,MCP客户端准备就绪")
logger.bind(tag=TAG).info("所有工具已获取,MCP客户端准备就绪")
# 刷新工具缓存,确保MCP工具被包含在函数列表中
if hasattr(conn, "func_handler") and conn.func_handler:
conn.func_handler.tool_manager.refresh_tools()
conn.func_handler.current_support_functions()
return
# Handle method calls (requests from the client)
elif "method" in payload:
method = payload["method"]
conn.logger.bind(tag=TAG).info(f"收到MCP客户端请求: {method}")
logger.bind(tag=TAG).info(f"收到MCP客户端请求: {method}")
elif "error" in payload:
error_data = payload["error"]
error_msg = error_data.get("message", "未知错误")
conn.logger.bind(tag=TAG).error(f"收到MCP错误响应: {error_msg}")
logger.bind(tag=TAG).error(f"收到MCP错误响应: {error_msg}")
msg_id = int(payload.get("id", 0))
if msg_id in mcp_client.call_results:
@@ -216,9 +224,6 @@ async def handle_mcp_message(conn, mcp_client: MCPClient, payload: dict):
)
# --- Outgoing MCP Messages ---
async def send_mcp_initialize_message(conn):
"""发送MCP初始化消息"""
@@ -250,7 +255,7 @@ async def send_mcp_initialize_message(conn):
},
},
}
conn.logger.bind(tag=TAG).info("发送MCP初始化消息")
logger.bind(tag=TAG).info("发送MCP初始化消息")
await send_mcp_message(conn, payload)
@@ -261,7 +266,7 @@ async def send_mcp_tools_list_request(conn):
"id": 2, # mcpToolsListID
"method": "tools/list",
}
conn.logger.bind(tag=TAG).debug("发送MCP工具列表请求")
logger.bind(tag=TAG).debug("发送MCP工具列表请求")
await send_mcp_message(conn, payload)
@@ -273,7 +278,7 @@ async def send_mcp_tools_list_continue_request(conn, cursor: str):
"method": "tools/list",
"params": {"cursor": cursor},
}
conn.logger.bind(tag=TAG).info(f"发送带cursor的MCP工具列表请求: {cursor}")
logger.bind(tag=TAG).info(f"发送带cursor的MCP工具列表请求: {cursor}")
await send_mcp_message(conn, payload)
@@ -307,8 +312,6 @@ async def call_mcp_tool(
# 如果解析失败,尝试合并多个JSON对象
try:
# 使用正则表达式匹配所有JSON对象
import re
json_objects = re.findall(r"\{[^{}]*\}", args)
if len(json_objects) > 1:
# 合并所有JSON对象
@@ -327,7 +330,7 @@ async def call_mcp_tool(
else:
raise ValueError(f"参数JSON解析失败: {args}")
except Exception as e:
conn.logger.bind(tag=TAG).error(
logger.bind(tag=TAG).error(
f"参数JSON解析失败: {str(e)}, 原始参数: {args}"
)
raise ValueError(f"参数JSON解析失败: {str(e)}")
@@ -353,15 +356,13 @@ async def call_mcp_tool(
"params": {"name": actual_name, "arguments": arguments},
}
conn.logger.bind(tag=TAG).info(
f"发送客户端mcp工具调用请求: {actual_name},参数: {args}"
)
logger.bind(tag=TAG).info(f"发送客户端mcp工具调用请求: {actual_name},参数: {args}")
await send_mcp_message(conn, payload)
try:
# Wait for response or timeout
raw_result = await asyncio.wait_for(result_future, timeout=timeout)
conn.logger.bind(tag=TAG).info(
logger.bind(tag=TAG).info(
f"客户端mcp工具调用 {actual_name} 成功,原始结果: {raw_result}"
)
@@ -0,0 +1,21 @@
"""MCP接入点工具模块"""
from .mcp_endpoint_executor import MCPEndpointExecutor
from .mcp_endpoint_client import MCPEndpointClient
from .mcp_endpoint_handler import (
connect_mcp_endpoint,
send_mcp_endpoint_initialize,
send_mcp_endpoint_notification,
send_mcp_endpoint_tools_list,
call_mcp_endpoint_tool,
)
__all__ = [
"MCPEndpointExecutor",
"MCPEndpointClient",
"connect_mcp_endpoint",
"send_mcp_endpoint_initialize",
"send_mcp_endpoint_notification",
"send_mcp_endpoint_tools_list",
"call_mcp_endpoint_tool",
]
@@ -0,0 +1,112 @@
"""MCP接入点客户端定义"""
import asyncio
from concurrent.futures import Future
from core.utils.util import sanitize_tool_name
from config.logger import setup_logging
TAG = __name__
logger = setup_logging()
class MCPEndpointClient:
"""MCP接入点客户端,用于管理MCP接入点状态和工具"""
def __init__(self, conn=None):
self.conn = conn
self.tools = {} # sanitized_name -> tool_data
self.name_mapping = {}
self.ready = False
self.call_results = {} # To store Futures for tool call responses
self.next_id = 1
self.lock = asyncio.Lock()
self._cached_available_tools = None # Cache for get_available_tools
self.websocket = None # WebSocket连接
def has_tool(self, name: str) -> bool:
return name in self.tools
def get_available_tools(self) -> list:
# Check if the cache is valid
if self._cached_available_tools is not None:
return self._cached_available_tools
# If cache is not valid, regenerate the list
result = []
for tool_name, tool_data in self.tools.items():
function_def = {
"name": tool_name,
"description": tool_data["description"],
"parameters": {
"type": tool_data["inputSchema"].get("type", "object"),
"properties": tool_data["inputSchema"].get("properties", {}),
"required": tool_data["inputSchema"].get("required", []),
},
}
result.append({"type": "function", "function": function_def})
self._cached_available_tools = result # Store the generated list in cache
return result
async def is_ready(self) -> bool:
async with self.lock:
return self.ready
async def set_ready(self, status: bool):
async with self.lock:
self.ready = status
async def add_tool(self, tool_data: dict):
async with self.lock:
sanitized_name = sanitize_tool_name(tool_data["name"])
self.tools[sanitized_name] = tool_data
self.name_mapping[sanitized_name] = tool_data["name"]
self._cached_available_tools = (
None # Invalidate the cache when a tool is added
)
async def get_next_id(self) -> int:
async with self.lock:
current_id = self.next_id
self.next_id += 1
return current_id
async def register_call_result_future(self, id: int, future: Future):
async with self.lock:
self.call_results[id] = future
async def resolve_call_result(self, id: int, result: any):
async with self.lock:
if id in self.call_results:
future = self.call_results.pop(id)
if not future.done():
future.set_result(result)
async def reject_call_result(self, id: int, exception: Exception):
async with self.lock:
if id in self.call_results:
future = self.call_results.pop(id)
if not future.done():
future.set_exception(exception)
async def cleanup_call_result(self, id: int):
async with self.lock:
if id in self.call_results:
self.call_results.pop(id)
def set_websocket(self, websocket):
"""设置WebSocket连接"""
self.websocket = websocket
async def send_message(self, message: str):
"""发送消息到MCP接入点"""
if self.websocket:
await self.websocket.send(message)
else:
raise RuntimeError("WebSocket连接未建立")
async def close(self):
"""关闭WebSocket连接"""
if self.websocket:
await self.websocket.close()
self.websocket = None
@@ -0,0 +1,97 @@
"""MCP接入点工具执行器"""
from typing import Dict, Any
from ..base import ToolType, ToolDefinition, ToolExecutor
from plugins_func.register import Action, ActionResponse
from .mcp_endpoint_handler import call_mcp_endpoint_tool
class MCPEndpointExecutor(ToolExecutor):
"""MCP接入点工具执行器"""
def __init__(self, conn):
self.conn = conn
async def execute(
self, conn, tool_name: str, arguments: Dict[str, Any]
) -> ActionResponse:
"""执行MCP接入点工具"""
if not hasattr(conn, "mcp_endpoint_client") or not conn.mcp_endpoint_client:
return ActionResponse(
action=Action.ERROR,
response="MCP接入点客户端未初始化",
)
if not await conn.mcp_endpoint_client.is_ready():
return ActionResponse(
action=Action.ERROR,
response="MCP接入点客户端未准备就绪",
)
try:
# 转换参数为JSON字符串
import json
args_str = json.dumps(arguments) if arguments else "{}"
# 调用MCP接入点工具
result = await call_mcp_endpoint_tool(
conn.mcp_endpoint_client, tool_name, args_str
)
resultJson = None
if isinstance(result, str):
try:
resultJson = json.loads(result)
except Exception as e:
pass
# 视觉大模型不经过二次LLM处理
if (
resultJson is not None
and isinstance(resultJson, dict)
and "action" in resultJson
):
return ActionResponse(
action=Action[resultJson["action"]],
response=resultJson.get("response", ""),
)
return ActionResponse(action=Action.REQLLM, result=str(result))
except ValueError as e:
return ActionResponse(action=Action.NOTFOUND, response=str(e))
except Exception as e:
return ActionResponse(action=Action.ERROR, response=str(e))
def get_tools(self) -> Dict[str, ToolDefinition]:
"""获取所有MCP接入点工具"""
if (
not hasattr(self.conn, "mcp_endpoint_client")
or not self.conn.mcp_endpoint_client
):
return {}
tools = {}
mcp_tools = self.conn.mcp_endpoint_client.get_available_tools()
for tool in mcp_tools:
func_def = tool.get("function", {})
tool_name = func_def.get("name", "")
if tool_name:
tools[tool_name] = ToolDefinition(
name=tool_name, description=tool, tool_type=ToolType.MCP_ENDPOINT
)
return tools
def has_tool(self, tool_name: str) -> bool:
"""检查是否有指定的MCP接入点工具"""
if (
not hasattr(self.conn, "mcp_endpoint_client")
or not self.conn.mcp_endpoint_client
):
return False
return self.conn.mcp_endpoint_client.has_tool(tool_name)
@@ -0,0 +1,390 @@
"""MCP接入点处理器"""
import json
import asyncio
import re
import websockets
from config.logger import setup_logging
from .mcp_endpoint_client import MCPEndpointClient
TAG = __name__
logger = setup_logging()
async def connect_mcp_endpoint(mcp_endpoint_url: str, conn=None) -> MCPEndpointClient:
"""连接到MCP接入点"""
if not mcp_endpoint_url or "你的" in mcp_endpoint_url or mcp_endpoint_url == "null":
return None
try:
websocket = await websockets.connect(mcp_endpoint_url)
mcp_client = MCPEndpointClient(conn)
mcp_client.set_websocket(websocket)
# 启动消息监听器
asyncio.create_task(_message_listener(mcp_client))
# 发送初始化消息
await send_mcp_endpoint_initialize(mcp_client)
# 发送初始化完成通知
await send_mcp_endpoint_notification(mcp_client, "notifications/initialized")
# 获取工具列表
await send_mcp_endpoint_tools_list(mcp_client)
logger.bind(tag=TAG).info("MCP接入点连接成功")
return mcp_client
except Exception as e:
logger.bind(tag=TAG).error(f"连接MCP接入点失败: {e}")
return None
async def _message_listener(mcp_client: MCPEndpointClient):
"""监听MCP接入点消息"""
try:
async for message in mcp_client.websocket:
await handle_mcp_endpoint_message(mcp_client, message)
except websockets.exceptions.ConnectionClosed:
logger.bind(tag=TAG).info("MCP接入点连接已关闭")
except Exception as e:
logger.bind(tag=TAG).error(f"MCP接入点消息监听器错误: {e}")
finally:
await mcp_client.set_ready(False)
async def handle_mcp_endpoint_message(mcp_client: MCPEndpointClient, message: str):
"""处理MCP接入点消息"""
try:
payload = json.loads(message)
logger.bind(tag=TAG).debug(f"收到MCP接入点消息: {payload}")
if not isinstance(payload, dict):
logger.bind(tag=TAG).error("MCP接入点消息格式错误")
return
# Handle result
if "result" in payload:
result = payload["result"]
# 安全地获取消息ID,如果为None则使用0
msg_id_raw = payload.get("id")
msg_id = int(msg_id_raw) if msg_id_raw is not None else 0
# Check for tool call response first
if msg_id in mcp_client.call_results:
logger.bind(tag=TAG).debug(
f"收到工具调用响应,ID: {msg_id}, 结果: {result}"
)
await mcp_client.resolve_call_result(msg_id, result)
return
if msg_id == 1: # mcpInitializeID
logger.bind(tag=TAG).debug("收到MCP接入点初始化响应")
if result is not None and isinstance(result, dict):
server_info = result.get("serverInfo")
if isinstance(server_info, dict):
name = server_info.get("name")
version = server_info.get("version")
logger.bind(tag=TAG).info(
f"MCP接入点服务器信息: name={name}, version={version}"
)
else:
logger.bind(tag=TAG).warning(
"MCP接入点初始化响应结果为空或格式错误"
)
return
elif msg_id == 2: # mcpToolsListID
logger.bind(tag=TAG).debug("收到MCP接入点工具列表响应")
if (
result is not None
and isinstance(result, dict)
and "tools" in result
):
tools_data = result["tools"]
if not isinstance(tools_data, list):
logger.bind(tag=TAG).error("工具列表格式错误")
return
logger.bind(tag=TAG).info(
f"MCP接入点支持的工具数量: {len(tools_data)}"
)
for i, tool in enumerate(tools_data):
if not isinstance(tool, dict):
continue
name = tool.get("name", "")
description = tool.get("description", "")
input_schema = {
"type": "object",
"properties": {},
"required": [],
}
if "inputSchema" in tool and isinstance(
tool["inputSchema"], dict
):
schema = tool["inputSchema"]
input_schema["type"] = schema.get("type", "object")
input_schema["properties"] = schema.get("properties", {})
input_schema["required"] = [
s
for s in schema.get("required", [])
if isinstance(s, str)
]
new_tool = {
"name": name,
"description": description,
"inputSchema": input_schema,
}
await mcp_client.add_tool(new_tool)
logger.bind(tag=TAG).debug(f"MCP接入点工具 #{i+1}: {name}")
# 替换所有工具描述中的工具名称
for tool_data in mcp_client.tools.values():
if "description" in tool_data:
description = tool_data["description"]
# 遍历所有工具名称进行替换
for (
sanitized_name,
original_name,
) in mcp_client.name_mapping.items():
description = description.replace(
original_name, sanitized_name
)
tool_data["description"] = description
next_cursor = (
result.get("nextCursor", "") if result is not None else ""
)
if next_cursor:
logger.bind(tag=TAG).info(
f"有更多工具,nextCursor: {next_cursor}"
)
await send_mcp_endpoint_tools_list_continue(
mcp_client, next_cursor
)
else:
await mcp_client.set_ready(True)
logger.bind(tag=TAG).info(
"所有MCP接入点工具已获取,客户端准备就绪"
)
# 刷新工具缓存,确保MCP接入点工具被包含在函数列表中
if (
hasattr(mcp_client, "conn")
and mcp_client.conn
and hasattr(mcp_client.conn, "func_handler")
and mcp_client.conn.func_handler
):
mcp_client.conn.func_handler.tool_manager.refresh_tools()
mcp_client.conn.func_handler.current_support_functions()
logger.bind(tag=TAG).info(
f"MCP接入点工具获取完成,共 {len(mcp_client.tools)} 个工具"
)
else:
logger.bind(tag=TAG).warning(
"MCP接入点工具列表响应结果为空或格式错误"
)
return
# Handle method calls (requests from the endpoint)
elif "method" in payload:
method = payload["method"]
logger.bind(tag=TAG).info(f"收到MCP接入点请求: {method}")
elif "error" in payload:
error_data = payload["error"]
error_msg = error_data.get("message", "未知错误")
logger.bind(tag=TAG).error(f"收到MCP接入点错误响应: {error_msg}")
# 安全地获取消息ID,如果为None则使用0
msg_id_raw = payload.get("id")
msg_id = int(msg_id_raw) if msg_id_raw is not None else 0
if msg_id in mcp_client.call_results:
await mcp_client.reject_call_result(
msg_id, Exception(f"MCP接入点错误: {error_msg}")
)
except json.JSONDecodeError as e:
logger.bind(tag=TAG).error(f"MCP接入点消息JSON解析失败: {e}")
except Exception as e:
logger.bind(tag=TAG).error(f"处理MCP接入点消息时出错: {e}")
import traceback
logger.bind(tag=TAG).error(f"错误详情: {traceback.format_exc()}")
async def send_mcp_endpoint_initialize(mcp_client: MCPEndpointClient):
"""发送MCP接入点初始化消息"""
payload = {
"jsonrpc": "2.0",
"id": 1, # mcpInitializeID
"method": "initialize",
"params": {
"protocolVersion": "2024-11-05",
"capabilities": {
"roots": {"listChanged": True},
"sampling": {},
},
"clientInfo": {
"name": "XiaozhiMCPEndpointClient",
"version": "1.0.0",
},
},
}
message = json.dumps(payload)
logger.bind(tag=TAG).info("发送MCP接入点初始化消息")
await mcp_client.send_message(message)
async def send_mcp_endpoint_notification(mcp_client: MCPEndpointClient, method: str):
"""发送MCP接入点通知消息"""
payload = {
"jsonrpc": "2.0",
"method": method,
"params": {},
}
message = json.dumps(payload)
logger.bind(tag=TAG).debug(f"发送MCP接入点通知: {method}")
await mcp_client.send_message(message)
async def send_mcp_endpoint_tools_list(mcp_client: MCPEndpointClient):
"""发送MCP接入点工具列表请求"""
payload = {
"jsonrpc": "2.0",
"id": 2, # mcpToolsListID
"method": "tools/list",
}
message = json.dumps(payload)
logger.bind(tag=TAG).debug("发送MCP接入点工具列表请求")
await mcp_client.send_message(message)
async def send_mcp_endpoint_tools_list_continue(
mcp_client: MCPEndpointClient, cursor: str
):
"""发送带有cursor的MCP接入点工具列表请求"""
payload = {
"jsonrpc": "2.0",
"id": 2, # mcpToolsListID (same ID for continuation)
"method": "tools/list",
"params": {"cursor": cursor},
}
message = json.dumps(payload)
logger.bind(tag=TAG).info(f"发送带cursor的MCP接入点工具列表请求: {cursor}")
await mcp_client.send_message(message)
async def call_mcp_endpoint_tool(
mcp_client: MCPEndpointClient, tool_name: str, args: str = "{}", timeout: int = 30
):
"""
调用指定的MCP接入点工具,并等待响应
"""
if not await mcp_client.is_ready():
raise RuntimeError("MCP接入点客户端尚未准备就绪")
if not mcp_client.has_tool(tool_name):
raise ValueError(f"工具 {tool_name} 不存在")
tool_call_id = await mcp_client.get_next_id()
result_future = asyncio.Future()
await mcp_client.register_call_result_future(tool_call_id, result_future)
# 处理参数
try:
if isinstance(args, str):
# 确保字符串是有效的JSON
if not args.strip():
arguments = {}
else:
try:
# 尝试直接解析
arguments = json.loads(args)
except json.JSONDecodeError:
# 如果解析失败,尝试合并多个JSON对象
try:
# 使用正则表达式匹配所有JSON对象
json_objects = re.findall(r"\{[^{}]*\}", args)
if len(json_objects) > 1:
# 合并所有JSON对象
merged_dict = {}
for json_str in json_objects:
try:
obj = json.loads(json_str)
if isinstance(obj, dict):
merged_dict.update(obj)
except json.JSONDecodeError:
continue
if merged_dict:
arguments = merged_dict
else:
raise ValueError(f"无法解析任何有效的JSON对象: {args}")
else:
raise ValueError(f"参数JSON解析失败: {args}")
except Exception as e:
logger.bind(tag=TAG).error(
f"参数JSON解析失败: {str(e)}, 原始参数: {args}"
)
raise ValueError(f"参数JSON解析失败: {str(e)}")
elif isinstance(args, dict):
arguments = args
else:
raise ValueError(f"参数类型错误,期望字符串或字典,实际类型: {type(args)}")
# 确保参数是字典类型
if not isinstance(arguments, dict):
raise ValueError(f"参数必须是字典类型,实际类型: {type(arguments)}")
except Exception as e:
if not isinstance(e, ValueError):
raise ValueError(f"参数处理失败: {str(e)}")
raise e
actual_name = mcp_client.name_mapping.get(tool_name, tool_name)
payload = {
"jsonrpc": "2.0",
"id": tool_call_id,
"method": "tools/call",
"params": {"name": actual_name, "arguments": arguments},
}
message = json.dumps(payload)
logger.bind(tag=TAG).info(f"发送MCP接入点工具调用请求: {actual_name},参数: {args}")
await mcp_client.send_message(message)
try:
# Wait for response or timeout
raw_result = await asyncio.wait_for(result_future, timeout=timeout)
logger.bind(tag=TAG).info(
f"MCP接入点工具调用 {actual_name} 成功,原始结果: {raw_result}"
)
if isinstance(raw_result, dict):
if raw_result.get("isError") is True:
error_msg = raw_result.get(
"error", "工具调用返回错误,但未提供具体错误信息"
)
raise RuntimeError(f"工具调用错误: {error_msg}")
content = raw_result.get("content")
if isinstance(content, list) and len(content) > 0:
if isinstance(content[0], dict) and "text" in content[0]:
# 直接返回文本内容,不进行JSON解析
return content[0]["text"]
# 如果结果不是预期的格式,将其转换为字符串
return str(raw_result)
except asyncio.TimeoutError:
await mcp_client.cleanup_call_result(tool_call_id)
raise TimeoutError("工具调用请求超时")
except Exception as e:
await mcp_client.cleanup_call_result(tool_call_id)
raise e
@@ -0,0 +1,7 @@
"""服务端MCP工具模块"""
from .mcp_manager import ServerMCPManager
from .mcp_executor import ServerMCPExecutor
from .mcp_client import ServerMCPClient
__all__ = ["ServerMCPManager", "ServerMCPExecutor", "ServerMCPClient"]
@@ -1,7 +1,12 @@
"""服务端MCP客户端"""
from __future__ import annotations
from datetime import timedelta
import asyncio, os, shutil, concurrent.futures
import asyncio
import os
import shutil
import concurrent.futures
from contextlib import AsyncExitStack
from typing import Optional, List, Dict, Any
@@ -14,8 +19,15 @@ from core.utils.util import sanitize_tool_name
TAG = __name__
class MCPClient:
class ServerMCPClient:
"""服务端MCP客户端,用于连接和管理MCP服务"""
def __init__(self, config: Dict[str, Any]):
"""初始化服务端MCP客户端
Args:
config: MCP服务配置字典
"""
self.logger = setup_logging()
self.config = config
@@ -24,21 +36,26 @@ class MCPClient:
self._shutdown_evt = asyncio.Event()
self.session: Optional[ClientSession] = None
self.tools: List = [] # original tool objects
self.tools: List = [] # 原始工具对象
self.tools_dict: Dict[str, Any] = {}
self.name_mapping: Dict[str, str] = {}
async def initialize(self):
"""初始化MCP客户端连接"""
if self._worker_task:
return
self._worker_task = asyncio.create_task(self._worker(), name="MCPClientWorker")
self._worker_task = asyncio.create_task(
self._worker(), name="ServerMCPClientWorker"
)
await self._ready_evt.wait()
self.logger.bind(tag=TAG).info(
f"Connected, tools = {[name for name in self.name_mapping.values()]}"
f"服务端MCP客户端已连接,可用工具: {[name for name in self.name_mapping.values()]}"
)
async def cleanup(self):
"""清理MCP客户端资源"""
if not self._worker_task:
return
@@ -46,14 +63,27 @@ class MCPClient:
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}")
self.logger.bind(tag=TAG).error(f"服务端MCP客户端关闭错误: {e}")
finally:
self._worker_task = None
def has_tool(self, name: str) -> bool:
"""检查是否包含指定工具
Args:
name: 工具名称
Returns:
bool: 是否包含该工具
"""
return name in self.tools_dict
def get_available_tools(self):
def get_available_tools(self) -> List[Dict[str, Any]]:
"""获取所有可用工具的定义
Returns:
List[Dict[str, Any]]: 工具定义列表
"""
return [
{
"type": "function",
@@ -66,9 +96,21 @@ class MCPClient:
for name, tool in self.tools_dict.items()
]
async def call_tool(self, name: str, args: dict):
async def call_tool(self, name: str, args: dict) -> Any:
"""调用指定工具
Args:
name: 工具名称
args: 工具参数
Returns:
Any: 工具执行结果
Raises:
RuntimeError: 客户端未初始化时抛出
"""
if not self.session:
raise RuntimeError("MCPClient not initialized")
raise RuntimeError("服务端MCP客户端未初始化")
real_name = self.name_mapping.get(name, name)
loop = self._worker_task.get_loop()
@@ -80,7 +122,29 @@ class MCPClient:
fut: concurrent.futures.Future = asyncio.run_coroutine_threadsafe(coro, loop)
return await asyncio.wrap_future(fut)
def is_connected(self) -> bool:
"""检查MCP客户端是否连接正常
Returns:
bool: 如果客户端已连接并正常工作返回True否则返回False
"""
# 检查工作任务是否存在
if self._worker_task is None:
return False
# 检查工作任务是否已经完成或取消
if self._worker_task.done():
return False
# 检查会话是否存在
if self.session is None:
return False
# 所有检查都通过,连接正常
return True
async def _worker(self):
"""MCP客户端工作协程"""
async with AsyncExitStack() as stack:
try:
# 建立 StdioClient
@@ -100,6 +164,7 @@ class MCPClient:
stdio_client(params)
)
read_stream, write_stream = stdio_r, stdio_w
# 建立SSEClient
elif "url" in self.config:
if "API_ACCESS_TOKEN" in self.config:
@@ -114,7 +179,7 @@ class MCPClient:
read_stream, write_stream = sse_r, sse_w
else:
raise ValueError("MCPClient config must include 'command' or 'url'")
raise ValueError("MCP客户端配置必须包含'command''url'")
self.session = await stack.enter_async_context(
ClientSession(
@@ -138,6 +203,6 @@ class MCPClient:
await self._shutdown_evt.wait()
except Exception as e:
self.logger.bind(tag=TAG).error(f"worker error: {e}")
self.logger.bind(tag=TAG).error(f"服务端MCP客户端工作协程错误: {e}")
self._ready_evt.set()
raise
@@ -0,0 +1,89 @@
"""服务端MCP工具执行器"""
from typing import Dict, Any, Optional
from ..base import ToolType, ToolDefinition, ToolExecutor
from plugins_func.register import Action, ActionResponse
from .mcp_manager import ServerMCPManager
class ServerMCPExecutor(ToolExecutor):
"""服务端MCP工具执行器"""
def __init__(self, conn):
self.conn = conn
self.mcp_manager: Optional[ServerMCPManager] = None
self._initialized = False
async def initialize(self):
"""初始化MCP管理器"""
if not self._initialized:
self.mcp_manager = ServerMCPManager(self.conn)
await self.mcp_manager.initialize_servers()
self._initialized = True
async def execute(
self, conn, tool_name: str, arguments: Dict[str, Any]
) -> ActionResponse:
"""执行服务端MCP工具"""
if not self._initialized or not self.mcp_manager:
return ActionResponse(
action=Action.ERROR,
response="MCP管理器未初始化",
)
try:
# 移除mcp_前缀(如果有)
actual_tool_name = tool_name
if tool_name.startswith("mcp_"):
actual_tool_name = tool_name[4:]
result = await self.mcp_manager.execute_tool(actual_tool_name, arguments)
return ActionResponse(action=Action.REQLLM, result=str(result))
except ValueError as e:
return ActionResponse(
action=Action.NOTFOUND,
response=str(e),
)
except Exception as e:
return ActionResponse(
action=Action.ERROR,
response=str(e),
)
def get_tools(self) -> Dict[str, ToolDefinition]:
"""获取所有服务端MCP工具"""
if not self._initialized or not self.mcp_manager:
return {}
tools = {}
mcp_tools = self.mcp_manager.get_all_tools()
for tool in mcp_tools:
func_def = tool.get("function", {})
tool_name = func_def.get("name", "")
if tool_name == "":
continue
tools[tool_name] = ToolDefinition(
name=tool_name, description=tool, tool_type=ToolType.SERVER_MCP
)
return tools
def has_tool(self, tool_name: str) -> bool:
"""检查是否有指定的服务端MCP工具"""
if not self._initialized or not self.mcp_manager:
return False
# 移除mcp_前缀(如果有)
actual_tool_name = tool_name
if tool_name.startswith("mcp_"):
actual_tool_name = tool_name[4:]
return self.mcp_manager.is_mcp_tool(actual_tool_name)
async def cleanup(self):
"""清理MCP连接"""
if self.mcp_manager:
await self.mcp_manager.cleanup_all()
@@ -0,0 +1,158 @@
"""服务端MCP管理器"""
import asyncio
import os
import json
from typing import Dict, Any, List
from config.config_loader import get_project_dir
from config.logger import setup_logging
from .mcp_client import ServerMCPClient
TAG = __name__
logger = setup_logging()
class ServerMCPManager:
"""管理多个服务端MCP服务的集中管理器"""
def __init__(self, conn) -> None:
"""初始化MCP管理器"""
self.conn = conn
self.config_path = get_project_dir() + "data/.mcp_server_settings.json"
if not os.path.exists(self.config_path):
self.config_path = ""
logger.bind(tag=TAG).warning(
f"请检查mcp服务配置文件:data/.mcp_server_settings.json"
)
self.clients: Dict[str, ServerMCPClient] = {}
self.tools = []
def load_config(self) -> Dict[str, Any]:
"""加载MCP服务配置"""
if len(self.config_path) == 0:
return {}
try:
with open(self.config_path, "r", encoding="utf-8") as f:
config = json.load(f)
return config.get("mcpServers", {})
except Exception as e:
logger.bind(tag=TAG).error(
f"Error loading MCP config from {self.config_path}: {e}"
)
return {}
async def initialize_servers(self) -> None:
"""初始化所有MCP服务"""
config = self.load_config()
for name, srv_config in config.items():
if not srv_config.get("command") and not srv_config.get("url"):
logger.bind(tag=TAG).warning(
f"Skipping server {name}: neither command nor url specified"
)
continue
try:
# 初始化服务端MCP客户端
logger.bind(tag=TAG).info(f"初始化服务端MCP客户端: {name}")
client = ServerMCPClient(srv_config)
await client.initialize()
self.clients[name] = client
client_tools = client.get_available_tools()
self.tools.extend(client_tools)
except Exception as e:
logger.bind(tag=TAG).error(
f"Failed to initialize MCP server {name}: {e}"
)
# 输出当前支持的服务端MCP工具列表
if hasattr(self.conn, "func_handler") and self.conn.func_handler:
self.conn.func_handler.current_support_functions()
def get_all_tools(self) -> List[Dict[str, Any]]:
"""获取所有服务的工具function定义"""
return self.tools
def is_mcp_tool(self, tool_name: str) -> bool:
"""检查是否是MCP工具"""
for tool in self.tools:
if (
tool.get("function") is not None
and tool["function"].get("name") == tool_name
):
return True
return False
async def execute_tool(self, tool_name: str, arguments: Dict[str, Any]) -> Any:
"""执行工具调用,失败时会尝试重新连接"""
logger.bind(tag=TAG).info(f"执行服务端MCP工具 {tool_name},参数: {arguments}")
max_retries = 3 # 最大重试次数
retry_interval = 2 # 重试间隔(秒)
# 找到对应的客户端
client_name = None
target_client = None
for name, client in self.clients.items():
if client.has_tool(tool_name):
client_name = name
target_client = client
break
if not target_client:
raise ValueError(f"工具 {tool_name} 在任意MCP服务中未找到")
# 带重试机制的工具调用
for attempt in range(max_retries):
try:
return await target_client.call_tool(tool_name, arguments)
except Exception as e:
# 最后一次尝试失败时直接抛出异常
if attempt == max_retries - 1:
raise
logger.bind(tag=TAG).warning(
f"执行工具 {tool_name} 失败 (尝试 {attempt+1}/{max_retries}): {e}"
)
# 尝试重新连接
logger.bind(tag=TAG).info(
f"重试前尝试重新连接 MCP 客户端 {client_name}"
)
try:
# 关闭旧的连接
await target_client.cleanup()
# 重新初始化客户端
config = self.load_config()
if client_name in config:
client = ServerMCPClient(config[client_name])
await client.initialize()
self.clients[client_name] = client
target_client = client
logger.bind(tag=TAG).info(
f"成功重新连接 MCP 客户端: {client_name}"
)
else:
logger.bind(tag=TAG).error(
f"Cannot reconnect MCP client {client_name}: config not found"
)
except Exception as reconnect_error:
logger.bind(tag=TAG).error(
f"Failed to reconnect MCP client {client_name}: {reconnect_error}"
)
# 等待一段时间再重试
await asyncio.sleep(retry_interval)
async def cleanup_all(self) -> None:
"""关闭所有 MCP客户端"""
for name, client in list(self.clients.items()):
try:
if hasattr(client, "cleanup"):
await asyncio.wait_for(client.cleanup(), timeout=20)
logger.bind(tag=TAG).info(f"服务端MCP客户端已关闭: {name}")
except (asyncio.TimeoutError, Exception) as e:
logger.bind(tag=TAG).error(f"关闭服务端MCP客户端 {name} 时出错: {e}")
self.clients.clear()
@@ -0,0 +1,5 @@
"""服务端插件工具模块"""
from .plugin_executor import ServerPluginExecutor
__all__ = ["ServerPluginExecutor"]
@@ -0,0 +1,84 @@
"""服务端插件工具执行器"""
from typing import Dict, Any
from ..base import ToolType, ToolDefinition, ToolExecutor
from plugins_func.register import all_function_registry, Action, ActionResponse
class ServerPluginExecutor(ToolExecutor):
"""服务端插件工具执行器"""
def __init__(self, conn):
self.conn = conn
self.config = conn.config
async def execute(
self, conn, tool_name: str, arguments: Dict[str, Any]
) -> ActionResponse:
"""执行服务端插件工具"""
func_item = all_function_registry.get(tool_name)
if not func_item:
return ActionResponse(
action=Action.NOTFOUND, response=f"插件函数 {tool_name} 不存在"
)
try:
# 根据工具类型决定如何调用
if hasattr(func_item, "type"):
func_type = func_item.type
if func_type.code in [4, 5]: # SYSTEM_CTL, IOT_CTL (需要conn参数)
result = func_item.func(conn, **arguments)
elif func_type.code == 2: # WAIT
result = func_item.func(**arguments)
elif func_type.code == 3: # CHANGE_SYS_PROMPT
result = func_item.func(conn, **arguments)
else:
result = func_item.func(**arguments)
else:
# 默认不传conn参数
result = func_item.func(**arguments)
return result
except Exception as e:
return ActionResponse(
action=Action.ERROR,
response=str(e),
)
def get_tools(self) -> Dict[str, ToolDefinition]:
"""获取所有注册的服务端插件工具"""
tools = {}
# 获取必要的函数
necessary_functions = ["handle_exit_intent", "get_time", "get_lunar"]
# 获取配置中的函数
config_functions = self.config["Intent"][
self.config["selected_module"]["Intent"]
].get("functions", [])
# 转换为列表
if not isinstance(config_functions, list):
try:
config_functions = list(config_functions)
except TypeError:
config_functions = []
# 合并所有需要的函数
all_required_functions = list(set(necessary_functions + config_functions))
for func_name in all_required_functions:
func_item = all_function_registry.get(func_name)
if func_item:
tools[func_name] = ToolDefinition(
name=func_name,
description=func_item.description,
tool_type=ToolType.SERVER_PLUGIN,
)
return tools
def has_tool(self, tool_name: str) -> bool:
"""检查是否有指定的服务端插件工具"""
return tool_name in all_function_registry
@@ -0,0 +1,235 @@
"""统一工具处理器"""
import json
from typing import Dict, List, Any, Optional
from config.logger import setup_logging
from plugins_func.loadplugins import auto_import_modules
from .base import ToolType
from plugins_func.register import Action, ActionResponse
from .unified_tool_manager import ToolManager
from .server_plugins import ServerPluginExecutor
from .server_mcp import ServerMCPExecutor
from .device_iot import DeviceIoTExecutor
from .device_mcp import DeviceMCPExecutor
from .mcp_endpoint import MCPEndpointExecutor
class UnifiedToolHandler:
"""统一工具处理器"""
def __init__(self, conn):
self.conn = conn
self.config = conn.config
self.logger = setup_logging()
# 创建工具管理器
self.tool_manager = ToolManager(conn)
# 创建各类执行器
self.server_plugin_executor = ServerPluginExecutor(conn)
self.server_mcp_executor = ServerMCPExecutor(conn)
self.device_iot_executor = DeviceIoTExecutor(conn)
self.device_mcp_executor = DeviceMCPExecutor(conn)
self.mcp_endpoint_executor = MCPEndpointExecutor(conn)
# 注册执行器
self.tool_manager.register_executor(
ToolType.SERVER_PLUGIN, self.server_plugin_executor
)
self.tool_manager.register_executor(
ToolType.SERVER_MCP, self.server_mcp_executor
)
self.tool_manager.register_executor(
ToolType.DEVICE_IOT, self.device_iot_executor
)
self.tool_manager.register_executor(
ToolType.DEVICE_MCP, self.device_mcp_executor
)
self.tool_manager.register_executor(
ToolType.MCP_ENDPOINT, self.mcp_endpoint_executor
)
# 初始化标志
self.finish_init = False
async def _initialize(self):
"""异步初始化"""
try:
# 自动导入插件模块
auto_import_modules("plugins_func.functions")
# 初始化服务端MCP
await self.server_mcp_executor.initialize()
# 初始化MCP接入点
await self._initialize_mcp_endpoint()
# 初始化Home Assistant(如果需要)
self._initialize_home_assistant()
self.finish_init = True
self.logger.info("统一工具处理器初始化完成")
# 输出当前支持的所有工具列表
self.current_support_functions()
except Exception as e:
self.logger.error(f"统一工具处理器初始化失败: {e}")
async def _initialize_mcp_endpoint(self):
"""初始化MCP接入点"""
try:
from .mcp_endpoint import connect_mcp_endpoint
# 从配置中获取MCP接入点URL
mcp_endpoint_url = self.config.get("mcp_endpoint", "")
if (
mcp_endpoint_url
and "你的" not in mcp_endpoint_url
and mcp_endpoint_url != "null"
):
self.logger.info(f"正在初始化MCP接入点: {mcp_endpoint_url}")
mcp_endpoint_client = await connect_mcp_endpoint(
mcp_endpoint_url, self.conn
)
if mcp_endpoint_client:
# 将MCP接入点客户端保存到连接对象中
self.conn.mcp_endpoint_client = mcp_endpoint_client
self.logger.info("MCP接入点初始化成功")
else:
self.logger.warning("MCP接入点初始化失败")
except Exception as e:
self.logger.error(f"初始化MCP接入点失败: {e}")
def _initialize_home_assistant(self):
"""初始化Home Assistant提示词"""
try:
from plugins_func.functions.hass_init import append_devices_to_prompt
append_devices_to_prompt(self.conn)
except ImportError:
pass # 忽略导入错误
except Exception as e:
self.logger.error(f"初始化Home Assistant失败: {e}")
def get_functions(self) -> List[Dict[str, Any]]:
"""获取所有工具的函数描述"""
return self.tool_manager.get_function_descriptions()
def current_support_functions(self) -> List[str]:
"""获取当前支持的函数名称列表"""
func_names = self.tool_manager.get_supported_tool_names()
self.logger.info(f"当前支持的函数列表: {func_names}")
return func_names
def upload_functions_desc(self):
"""刷新函数描述列表"""
self.tool_manager.refresh_tools()
self.logger.info("函数描述列表已刷新")
def has_tool(self, tool_name: str) -> bool:
"""检查是否有指定工具"""
return self.tool_manager.has_tool(tool_name)
async def handle_llm_function_call(
self, conn, function_call_data: Dict[str, Any]
) -> Optional[ActionResponse]:
"""处理LLM函数调用"""
try:
# 处理多函数调用
if "function_calls" in function_call_data:
responses = []
for call in function_call_data["function_calls"]:
result = await self.tool_manager.execute_tool(
call["name"], call.get("arguments", {})
)
responses.append(result)
return self._combine_responses(responses)
# 处理单函数调用
function_name = function_call_data["name"]
arguments = function_call_data.get("arguments", {})
# 如果arguments是字符串,尝试解析为JSON
if isinstance(arguments, str):
try:
arguments = json.loads(arguments) if arguments else {}
except json.JSONDecodeError:
self.logger.error(f"无法解析函数参数: {arguments}")
return ActionResponse(
action=Action.ERROR,
response="无法解析函数参数",
)
self.logger.debug(f"调用函数: {function_name}, 参数: {arguments}")
# 执行工具调用
result = await self.tool_manager.execute_tool(function_name, arguments)
return result
except Exception as e:
self.logger.error(f"处理function call错误: {e}")
return ActionResponse(action=Action.ERROR, response=str(e))
def _combine_responses(self, responses: List[ActionResponse]) -> ActionResponse:
"""合并多个函数调用的响应"""
if not responses:
return ActionResponse(action=Action.NONE, response="无响应")
# 如果有任何错误,返回第一个错误
for response in responses:
if response.action == Action.ERROR:
return response
# 合并所有成功的响应
contents = []
responses_text = []
for response in responses:
if response.content:
contents.append(response.content)
if response.response:
responses_text.append(response.response)
# 确定最终的动作类型
final_action = Action.RESPONSE
for response in responses:
if response.action == Action.REQLLM:
final_action = Action.REQLLM
break
return ActionResponse(
action=final_action,
result="; ".join(contents) if contents else None,
response="; ".join(responses_text) if responses_text else None,
)
async def register_iot_tools(self, descriptors: List[Dict[str, Any]]):
"""注册IoT设备工具"""
self.device_iot_executor.register_iot_tools(descriptors)
self.tool_manager.refresh_tools()
self.logger.info(f"注册了{len(descriptors)}个IoT设备的工具")
def get_tool_statistics(self) -> Dict[str, int]:
"""获取工具统计信息"""
return self.tool_manager.get_tool_statistics()
async def cleanup(self):
"""清理资源"""
try:
await self.server_mcp_executor.cleanup()
# 清理MCP接入点连接
if (
hasattr(self.conn, "mcp_endpoint_client")
and self.conn.mcp_endpoint_client
):
await self.conn.mcp_endpoint_client.close()
self.logger.info("工具处理器清理完成")
except Exception as e:
self.logger.error(f"工具处理器清理失败: {e}")
@@ -0,0 +1,124 @@
"""统一工具管理器"""
from typing import Dict, List, Optional, Any
from config.logger import setup_logging
from plugins_func.register import Action, ActionResponse
from .base import ToolType, ToolDefinition, ToolExecutor
class ToolManager:
"""统一工具管理器,管理所有类型的工具"""
def __init__(self, conn):
self.conn = conn
self.logger = setup_logging()
self.executors: Dict[ToolType, ToolExecutor] = {}
self._cached_tools: Optional[Dict[str, ToolDefinition]] = None
self._cached_function_descriptions: Optional[List[Dict[str, Any]]] = None
def register_executor(self, tool_type: ToolType, executor: ToolExecutor):
"""注册工具执行器"""
self.executors[tool_type] = executor
self._invalidate_cache()
self.logger.info(f"注册工具执行器: {tool_type.value}")
def _invalidate_cache(self):
"""使缓存失效"""
self._cached_tools = None
self._cached_function_descriptions = None
def get_all_tools(self) -> Dict[str, ToolDefinition]:
"""获取所有工具定义"""
if self._cached_tools is not None:
return self._cached_tools
all_tools = {}
for tool_type, executor in self.executors.items():
try:
tools = executor.get_tools()
for name, definition in tools.items():
if name in all_tools:
self.logger.warning(f"工具名称冲突: {name}")
all_tools[name] = definition
except Exception as e:
self.logger.error(f"获取{tool_type.value}工具时出错: {e}")
self._cached_tools = all_tools
return all_tools
def get_function_descriptions(self) -> List[Dict[str, Any]]:
"""获取所有工具的函数描述(OpenAI格式)"""
if self._cached_function_descriptions is not None:
return self._cached_function_descriptions
descriptions = []
tools = self.get_all_tools()
for tool_definition in tools.values():
descriptions.append(tool_definition.description)
self._cached_function_descriptions = descriptions
return descriptions
def has_tool(self, tool_name: str) -> bool:
"""检查是否存在指定工具"""
tools = self.get_all_tools()
return tool_name in tools
def get_tool_type(self, tool_name: str) -> Optional[ToolType]:
"""获取工具类型"""
tools = self.get_all_tools()
tool_def = tools.get(tool_name)
return tool_def.tool_type if tool_def else None
async def execute_tool(
self, tool_name: str, arguments: Dict[str, Any]
) -> ActionResponse:
"""执行工具调用"""
try:
# 查找工具类型
tool_type = self.get_tool_type(tool_name)
if not tool_type:
return ActionResponse(
action=Action.NOTFOUND,
response=f"工具 {tool_name} 不存在",
)
# 获取对应的执行器
executor = self.executors.get(tool_type)
if not executor:
return ActionResponse(
action=Action.ERROR,
response=f"工具类型 {tool_type.value} 的执行器未注册",
)
# 执行工具
self.logger.info(f"执行工具: {tool_name},参数: {arguments}")
result = await executor.execute(self.conn, tool_name, arguments)
self.logger.debug(f"工具执行结果: {result}")
return result
except Exception as e:
self.logger.error(f"执行工具 {tool_name} 时出错: {e}")
return ActionResponse(action=Action.ERROR, response=str(e))
def get_supported_tool_names(self) -> List[str]:
"""获取所有支持的工具名称"""
tools = self.get_all_tools()
return list(tools.keys())
def refresh_tools(self):
"""刷新工具缓存"""
self._invalidate_cache()
self.logger.info("工具缓存已刷新")
def get_tool_statistics(self) -> Dict[str, int]:
"""获取工具统计信息"""
stats = {}
for tool_type, executor in self.executors.items():
try:
tools = executor.get_tools()
stats[tool_type.value] = len(tools)
except Exception as e:
self.logger.error(f"获取{tool_type.value}工具统计时出错: {e}")
stats[tool_type.value] = 0
return stats
@@ -43,7 +43,6 @@ class TTSProviderBase(ABC):
self.tts_text_buff = []
self.punctuations = (
"",
".",
"",
"?",
"",
@@ -59,7 +58,6 @@ class TTSProviderBase(ABC):
"",
",",
"",
".",
"",
"?",
"",
@@ -171,7 +169,7 @@ class TTSProviderBase(ABC):
)
)
# 对于单句的文本,进行分段处理
segments = re.split(r'([。!?!?;\n])', content_detail)
segments = re.split(r"([。!?!?;\n])", content_detail)
for seg in segments:
self.tts_text_queue.put(
TTSMessageDTO(
@@ -85,8 +85,12 @@ class TTSProvider(TTSProviderBase):
self.reference_id = (
None if not config.get("reference_id") else config.get("reference_id")
)
self.reference_audio = parse_string_to_list(config.get("reference_audio"))
self.reference_text = parse_string_to_list(config.get("reference_text"))
self.reference_audio = parse_string_to_list(
config.get('ref_audio')if config.get('ref_audio') else config.get("reference_audio")
)
self.reference_text = parse_string_to_list(
config.get('ref_text')if config.get('ref_text') else config.get("reference_text")
)
self.format = config.get("response_format", "wav")
self.audio_file_type = config.get("response_format", "wav")
self.api_key = config.get("api_key", "YOUR_API_KEY")
@@ -12,8 +12,8 @@ class TTSProvider(TTSProviderBase):
super().__init__(config, delete_audio_file)
self.url = config.get("url")
self.text_lang = config.get("text_lang", "zh")
self.ref_audio_path = config.get("ref_audio_path")
self.prompt_text = config.get("prompt_text")
self.ref_audio_path = config.get('ref_audio') if config.get('ref_audio') else config.get("ref_audio_path")
self.prompt_text = config.get('ref_text') if config.get('ref_text') else config.get("prompt_text")
self.prompt_lang = config.get("prompt_lang", "zh")
# 处理空字符串的情况
@@ -11,8 +11,8 @@ class TTSProvider(TTSProviderBase):
def __init__(self, config, delete_audio_file):
super().__init__(config, delete_audio_file)
self.url = config.get("url")
self.refer_wav_path = config.get("refer_wav_path")
self.prompt_text = config.get("prompt_text")
self.refer_wav_path = config.get('ref_audio')if config.get('ref_audio') else config.get("refer_wav_path")
self.prompt_text = config.get('ref_text')if config.get('ref_text') else config.get("prompt_text")
self.prompt_language = config.get("prompt_language")
self.text_language = config.get("text_language", "audo")
@@ -12,6 +12,8 @@ from core.utils.util import check_model_key
from core.providers.tts.base import TTSProviderBase
from core.handle.abortHandle import handleAbortMessage
from core.providers.tts.dto.dto import SentenceType, ContentType, InterfaceType
from asyncio import Task
TAG = __name__
logger = setup_logging()
@@ -141,6 +143,7 @@ class TTSProvider(TTSProviderBase):
super().__init__(config, delete_audio_file)
self.ws = None
self.interface_type = InterfaceType.DUAL_STREAM
self._monitor_task = None # 监听任务引用
self.appId = config.get("appid")
self.access_token = config.get("access_token")
self.cluster = config.get("cluster")
@@ -270,8 +273,7 @@ class TTSProvider(TTSProviderBase):
try:
# 建立新连接
if self.ws is None:
await handleAbortMessage(self.conn)
logger.bind(tag=TAG).error(f"WebSocket连接不存在,终止发送文本")
logger.bind(tag=TAG).warning(f"WebSocket连接不存在,终止发送文本")
return
# 过滤Markdown
@@ -293,6 +295,25 @@ class TTSProvider(TTSProviderBase):
async def start_session(self, session_id):
logger.bind(tag=TAG).info(f"开始会话~~{session_id}")
try:
task = self._monitor_task
if (
task is not None
and isinstance(task, Task)
and not task.done()
):
logger.bind(tag=TAG).info("等待上一个监听任务结束...")
if self.ws is not None:
logger.bind(tag=TAG).info("强制关闭上一个WebSocket连接以唤醒监听任务...")
try:
await self.ws.close()
except Exception as e:
logger.bind(tag=TAG).warning(f"关闭上一个ws异常: {e}")
self.ws = None
try:
await asyncio.wait_for(task, timeout=8)
except Exception as e:
logger.bind(tag=TAG).warning(f"等待监听任务异常: {e}")
self._monitor_task = None
# 建立新连接
await self._ensure_connection()
@@ -463,6 +484,8 @@ class TTSProvider(TTSProviderBase):
except:
pass
self.ws = None
# 监听任务退出时清理引用
self._monitor_task = None
async def send_event(
self,
@@ -35,6 +35,10 @@ class VADProvider(VADProviderBase):
pcm_frame = self.decoder.decode(opus_packet, 960)
conn.client_audio_buffer.extend(pcm_frame) # 将新数据加入缓冲区
# 确保帧计数器存在
if not hasattr(conn, "client_voice_frame_count"):
conn.client_voice_frame_count = 0
# 处理缓冲区中的完整帧(每次处理512采样点)
client_have_voice = False
while len(conn.client_audio_buffer) >= 512 * 2:
@@ -50,18 +54,24 @@ class VADProvider(VADProviderBase):
# 检测语音活动
with torch.no_grad():
speech_prob = self.model(audio_tensor, 16000).item()
client_have_voice = speech_prob >= self.vad_threshold
is_voice = speech_prob >= self.vad_threshold
if is_voice:
conn.client_voice_frame_count += 1
else:
conn.client_voice_frame_count = 0
# 只有连续4帧检测到语音才认为有语音
client_have_voice = conn.client_voice_frame_count >= 4
# 如果之前有声音,但本次没有声音,且与上次有声音的时间差已经超过了静默阈值,则认为已经说完一句话
if conn.client_have_voice and not client_have_voice:
stop_duration = (
time.time() * 1000 - conn.client_have_voice_last_time
)
stop_duration = time.time() * 1000 - conn.last_activity_time
if stop_duration >= self.silence_threshold_ms:
conn.client_voice_stop = True
if client_have_voice:
conn.client_have_voice = True
conn.client_have_voice_last_time = time.time() * 1000
conn.last_activity_time = time.time() * 1000
return client_have_voice
except opuslib_next.OpusError as e:
+2 -1
View File
@@ -980,4 +980,5 @@ def is_valid_image_file(file_data: bytes) -> bool:
def sanitize_tool_name(name: str) -> str:
"""Sanitize tool names for OpenAI compatibility."""
return re.sub(r"[^a-zA-Z0-9_-]", "_", name)
# 支持中文、英文字母、数字、下划线和连字符
return re.sub(r"[^a-zA-Z0-9_\-\u4e00-\u9fff]", "_", name)
+3 -1
View File
@@ -48,11 +48,12 @@ services:
- SPRING_DATASOURCE_DRUID_USERNAME=root
- SPRING_DATASOURCE_DRUID_PASSWORD=123456
- SPRING_DATA_REDIS_HOST=xiaozhi-esp32-server-redis
- SPRING_DATA_REDIS_PASSWORD=
- SPRING_DATA_REDIS_PORT=6379
volumes:
# 配置文件目录
- ./uploadfile:/uploadfile
# 数据库模块
xiaozhi-esp32-server-db:
image: mysql:latest
container_name: xiaozhi-esp32-server-db
@@ -73,6 +74,7 @@ services:
- MYSQL_ROOT_PASSWORD=123456
- MYSQL_DATABASE=xiaozhi_esp32_server
- MYSQL_INITDB_ARGS="--character-set-server=utf8mb4 --collation-server=utf8mb4_unicode_ci"
# redis模块
xiaozhi-esp32-server-redis:
image: redis
expose:
@@ -8,30 +8,90 @@ from markitdown import MarkItDown
TAG = __name__
logger = setup_logging()
# 新闻来源字典,包含名称和对应的API ID
NEWS_SOURCES = {
"thepaper": "澎湃新闻",
"baidu": "百度热搜",
"cls-depth": "财联社",
CHANNEL_MAP = {
"V2EX": "v2ex-share",
"知乎": "zhihu",
"微博": "weibo",
"联合早报": "zaobao",
"酷安": "coolapk",
"MKTNews": "mktnews-flash",
"华尔街见闻": "wallstreetcn-quick",
"36氪": "36kr-quick",
"抖音": "douyin",
"虎扑": "hupu",
"百度贴吧": "tieba",
"今日头条": "toutiao",
"IT之家": "ithome",
"澎湃新闻": "thepaper",
"卫星通讯社": "sputniknewscn",
"参考消息": "cankaoxiaoxi",
"远景论坛": "pcbeta-windows11",
"财联社": "cls-depth",
"雪球": "xueqiu-hotstock",
"格隆汇": "gelonghui",
"法布财经": "fastbull-express",
"Solidot": "solidot",
"Hacker News": "hackernews",
"Product Hunt": "producthunt",
"Github": "github-trending-today",
"哔哩哔哩": "bilibili-hot-search",
"快手": "kuaishou",
"靠谱新闻": "kaopu",
"金十数据": "jin10",
"百度热搜": "baidu",
"牛客": "nowcoder",
"少数派": "sspai",
"稀土掘金": "juejin",
"凤凰网": "ifeng",
"虫部落": "chongbuluo-latest",
}
# 动态生成新闻源描述
def generate_news_sources_description():
sources_desc = []
for source_id, source_name in NEWS_SOURCES.items():
sources_desc.append(f"{source_name}({source_id})")
return "".join(sources_desc)
# 默认新闻来源字典,当配置中没有指定时使用
DEFAULT_NEWS_SOURCES = "澎湃新闻;百度热搜;财联社"
def get_news_sources_from_config(conn):
"""从配置中获取新闻源字符串"""
try:
# 尝试从插件配置中获取新闻源
if (
conn.config.get("plugins")
and conn.config["plugins"].get("get_news_from_newsnow")
and conn.config["plugins"]["get_news_from_newsnow"].get("news_sources")
):
# 获取配置的新闻源字符串
news_sources_config = conn.config["plugins"]["get_news_from_newsnow"][
"news_sources"
]
if isinstance(news_sources_config, str) and news_sources_config.strip():
logger.bind(tag=TAG).debug(f"使用配置的新闻源: {news_sources_config}")
return news_sources_config
else:
logger.bind(tag=TAG).warning("新闻源配置为空或格式错误,使用默认配置")
else:
logger.bind(tag=TAG).debug("未找到新闻源配置,使用默认配置")
return DEFAULT_NEWS_SOURCES
except Exception as e:
logger.bind(tag=TAG).error(f"获取新闻源配置失败: {e},使用默认配置")
return DEFAULT_NEWS_SOURCES
# 从CHANNEL_MAP获取所有可用的新闻源名称
available_sources = list(CHANNEL_MAP.keys())
example_sources_str = "".join(available_sources)
GET_NEWS_FROM_NEWSNOW_FUNCTION_DESC = {
"type": "function",
"function": {
"name": "get_news_from_newsnow",
"description": (
"获取最新新闻,随机选择一条新闻进行播报。"
f"用户可以选择不同的新闻源,{generate_news_sources_description()}等。"
"如果没有指定,默认从澎湃新闻获取。"
f"用户可以选择不同的新闻源,标准的名称是:{example_sources_str}"
"例如用户要求百度新闻,其实就是百度热搜。如果没有指定,默认从澎湃新闻获取。"
"用户可以要求获取详细内容,此时会获取新闻的详细内容。"
),
"parameters": {
@@ -39,7 +99,7 @@ GET_NEWS_FROM_NEWSNOW_FUNCTION_DESC = {
"properties": {
"source": {
"type": "string",
"description": f"新闻源,例如{generate_news_sources_description()}等。可选参数,如果不提供则使用默认新闻源",
"description": f"新闻源的标准中文名称,例如{example_sources_str}等。可选参数,如果不提供则使用默认新闻源",
},
"detail": {
"type": "boolean",
@@ -111,10 +171,13 @@ def fetch_news_detail(url):
ToolType.SYSTEM_CTL,
)
def get_news_from_newsnow(
conn, source: str = "thepaper", detail: bool = False, lang: str = "zh_CN"
conn, source: str = "澎湃新闻", detail: bool = False, lang: str = "zh_CN"
):
"""获取新闻并随机选择一条进行播报,或获取上一条新闻的详细内容"""
try:
# 获取当前配置的新闻源
news_sources = get_news_sources_from_config(conn)
# 如果detail为True,获取上一条新闻的详细内容
detail = str(detail).lower() == "true"
if detail:
@@ -132,7 +195,7 @@ def get_news_from_newsnow(
url = conn.last_newsnow_link.get("url")
title = conn.last_newsnow_link.get("title", "未知标题")
source_id = conn.last_newsnow_link.get("source_id", "thepaper")
source_name = NEWS_SOURCES.get(source_id, "未知来源")
source_name = CHANNEL_MAP.get(source_id, "未知来源")
if not url or url == "#":
return ActionResponse(
@@ -166,21 +229,32 @@ def get_news_from_newsnow(
return ActionResponse(Action.REQLLM, detail_report, None)
# 否则,获取新闻列表并随机选择一条
# 验证新闻源是否有效,如果无效则使用默认源
if source not in NEWS_SOURCES:
logger.bind(tag=TAG).warning(f"无效的新闻源: {source},使用默认源thepaper")
source = "thepaper"
# 将中文名称转换为英文ID
english_source_id = None
source_name = NEWS_SOURCES.get(source, "澎湃新闻")
logger.bind(tag=TAG).info(f"获取新闻: 新闻源={source}({source_name})")
# 检查输入的中文名称是否在配置的新闻源中
news_sources_list = [
name.strip() for name in news_sources.split(";") if name.strip()
]
if source in news_sources_list:
# 如果输入的中文名称在配置的新闻源中,在 CHANNEL_MAP 中查找对应的英文ID
english_source_id = CHANNEL_MAP.get(source)
# 如果找不到对应的英文ID,使用默认源
if not english_source_id:
logger.bind(tag=TAG).warning(f"无效的新闻源: {source},使用默认源澎湃新闻")
english_source_id = "thepaper"
source = "澎湃新闻"
logger.bind(tag=TAG).info(f"获取新闻: 新闻源={source}({english_source_id})")
# 获取新闻列表
news_items = fetch_news_from_api(conn, source)
news_items = fetch_news_from_api(conn, english_source_id)
if not news_items:
return ActionResponse(
Action.REQLLM,
f"抱歉,未能从{source_name}获取到新闻信息,请稍后再试或尝试其他新闻源。",
f"抱歉,未能从{source}获取到新闻信息,请稍后再试或尝试其他新闻源。",
None,
)
@@ -193,14 +267,14 @@ def get_news_from_newsnow(
conn.last_newsnow_link = {
"url": selected_news.get("url", "#"),
"title": selected_news.get("title", "未知标题"),
"source_id": source,
"source_id": english_source_id,
}
# 构建新闻报告
news_report = (
f"根据下列数据,用{lang}回应用户的新闻查询请求:\n\n"
f"新闻标题: {selected_news['title']}\n"
# f"新闻来源: {source_name}\n"
# f"新闻来源: {source}\n"
f"(请以自然、流畅的方式向用户播报这条新闻标题,"
f"提示用户可以要求获取详细内容,此时会获取新闻的详细内容。)"
)
@@ -31,7 +31,8 @@ hass_set_state_function_desc = {
"description": "只有在设置静音操作时才需要,设置静音的时候该值为true,取消静音时该值为false",
},
"rgb_color": {
"type": "list",
"type": "array",
"items": {"type": "integer"},
"description": "只有在设置颜色时需要,这里填目标颜色的rgb值",
},
},
+1 -1
View File
@@ -35,7 +35,7 @@ class Action(Enum):
class ActionResponse:
def __init__(self, action: Action, result, response):
def __init__(self, action: Action, result=None, response=None):
self.action = action # 动作类型
self.result = result # 动作产生的结果
self.response = response # 直接回复的内容
+1 -1
View File
@@ -30,7 +30,7 @@ baidu-aip==4.16.13
chardet==5.2.0
aioconsole==0.8.1
markitdown==0.1.1
mcp-proxy==0.6.0
mcp-proxy==0.8.0
PyJWT==2.8.0
psutil==7.0.0
portalocker==2.10.1