mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-23 15:43:54 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d5e0e8ab40 | ||
|
|
500b0ebde3 | ||
|
|
6f6c39da23 | ||
|
|
f8ccc7a92f | ||
|
|
95a678bb18 | ||
|
|
e43ce6ba3c | ||
|
|
7663b0b969 | ||
|
|
ba898b0874 | ||
|
|
04843010bc | ||
|
|
41c06efeee | ||
|
|
d0425fa31a | ||
|
|
08753b97df | ||
|
|
a86fef5f28 | ||
|
|
c5c2f59b31 | ||
|
|
69ee8ab438 | ||
|
|
4a24bc8c21 | ||
|
|
1d293244d3 | ||
|
|
6812c5ac40 | ||
|
|
90ad3e019f | ||
|
|
3071d2bdfa | ||
|
|
ec4694f859 | ||
|
|
da8435cfea | ||
|
|
0f36daf1fd | ||
|
|
2f78acaf4d | ||
|
|
9412a26bfc | ||
|
|
2348d9ceb3 | ||
|
|
61b4b6617b | ||
|
|
a8620f21a0 | ||
|
|
0631fe37ea | ||
|
|
21339aba42 | ||
|
|
660231978f | ||
|
|
dde203b158 | ||
|
|
0ee69fa2b2 | ||
|
|
8c1a8f7d55 | ||
|
|
ef03b665c9 | ||
|
|
b87906234d | ||
|
|
896b318c49 | ||
|
|
b3d441c385 | ||
|
|
96de670bfe | ||
|
|
bc2fc35cc9 | ||
|
|
993d5395b2 | ||
|
|
67f0b828ea | ||
|
|
879c1267b6 | ||
|
|
46c53e36a6 | ||
|
|
09f6605cfe | ||
|
|
12cee4027a | ||
|
|
979fea0d60 | ||
|
|
94553c54bd | ||
|
|
248db31c8b | ||
|
|
24bfa1ca15 | ||
|
|
1e97a8febc | ||
|
|
2742f2e1ff | ||
|
|
0a5ae70a7c | ||
|
|
22d53bd36e | ||
|
|
ebf68929ce | ||
|
|
ef3b373211 | ||
|
|
5012a51e1d | ||
|
|
fb1f476a3c | ||
|
|
1a31c8cd1d | ||
|
|
3e491c7d79 | ||
|
|
33345817d2 | ||
|
|
82973a685a | ||
|
|
e1d245068c | ||
|
|
b74517ffb0 | ||
|
|
231ae8dfa6 | ||
|
|
4ee9a47d41 | ||
|
|
d8bf5cdedf | ||
|
|
755e0edd44 | ||
|
|
178df82693 | ||
|
|
a88c3b2032 | ||
|
|
513396b905 | ||
|
|
9885d4758b |
@@ -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的接入状态,就会看到你扩展的功能列表了。
|
||||||
|
|
||||||
@@ -111,6 +111,11 @@ public interface Constant {
|
|||||||
*/
|
*/
|
||||||
String FILE_EXTENSION_SEG = ".";
|
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.6";
|
public static final String VERSION = "0.6.1";
|
||||||
|
|
||||||
/**
|
/**
|
||||||
* 无效固件URL
|
* 无效固件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();
|
||||||
|
}
|
||||||
|
|
||||||
|
}
|
||||||
+66
@@ -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() {
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
+25
@@ -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);
|
||||||
|
}
|
||||||
+180
@@ -0,0 +1,180 @@
|
|||||||
|
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.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.dto.McpJsonRpcResponse;
|
||||||
|
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(20, TimeUnit.SECONDS))) {
|
||||||
|
|
||||||
|
// 发送初始化通知
|
||||||
|
McpJsonRpcRequest initRequest = new McpJsonRpcRequest("notifications/initialized");
|
||||||
|
client.sendJson(initRequest);
|
||||||
|
|
||||||
|
// 等待 0.2 秒
|
||||||
|
Thread.sleep(200);
|
||||||
|
|
||||||
|
// 发送工具列表请求
|
||||||
|
McpJsonRpcRequest toolsRequest = new McpJsonRpcRequest("tools/list", null, 1);
|
||||||
|
client.sendJson(toolsRequest);
|
||||||
|
|
||||||
|
// 监听响应,直到收到包含 id=1 的响应
|
||||||
|
List<String> responses = client.listener(response -> {
|
||||||
|
try {
|
||||||
|
McpJsonRpcResponse jsonResponse = JsonUtils.parseObject(response, McpJsonRpcResponse.class);
|
||||||
|
return jsonResponse != null && Integer.valueOf(1).equals(jsonResponse.getId());
|
||||||
|
} catch (Exception e) {
|
||||||
|
log.warn("解析响应失败: {}", response, e);
|
||||||
|
return false;
|
||||||
|
}
|
||||||
|
});
|
||||||
|
|
||||||
|
// 处理响应
|
||||||
|
for (String response : responses) {
|
||||||
|
try {
|
||||||
|
McpJsonRpcResponse jsonResponse = JsonUtils.parseObject(response, McpJsonRpcResponse.class);
|
||||||
|
if (jsonResponse != null && Integer.valueOf(1).equals(jsonResponse.getId())
|
||||||
|
&& jsonResponse.getResult() != null && jsonResponse.getResult().getTools() != null) {
|
||||||
|
|
||||||
|
// 提取工具名称列表
|
||||||
|
return java.util.Arrays.stream(jsonResponse.getResult().getTools())
|
||||||
|
.map(McpJsonRpcResponse.McpTool::getName)
|
||||||
|
.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);
|
||||||
|
}
|
||||||
|
}
|
||||||
+3
@@ -56,6 +56,9 @@ public class AgentTemplateServiceImpl extends ServiceImpl<AgentTemplateDao, Agen
|
|||||||
wrapper.set("tts_model_id", modelId);
|
wrapper.set("tts_model_id", modelId);
|
||||||
wrapper.set("tts_voice_id", null);
|
wrapper.set("tts_voice_id", null);
|
||||||
break;
|
break;
|
||||||
|
case "VLLM":
|
||||||
|
wrapper.set("vllm_model_id", modelId);
|
||||||
|
break;
|
||||||
case "MEMORY":
|
case "MEMORY":
|
||||||
wrapper.set("mem_model_id", modelId);
|
wrapper.set("mem_model_id", modelId);
|
||||||
break;
|
break;
|
||||||
|
|||||||
+13
-1
@@ -1,6 +1,10 @@
|
|||||||
package xiaozhi.modules.config.service.impl;
|
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.apache.commons.lang3.StringUtils;
|
||||||
import org.springframework.stereotype.Service;
|
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.AgentEntity;
|
||||||
import xiaozhi.modules.agent.entity.AgentPluginMapping;
|
import xiaozhi.modules.agent.entity.AgentPluginMapping;
|
||||||
import xiaozhi.modules.agent.entity.AgentTemplateEntity;
|
import xiaozhi.modules.agent.entity.AgentTemplateEntity;
|
||||||
|
import xiaozhi.modules.agent.service.AgentMcpAccessPointService;
|
||||||
import xiaozhi.modules.agent.service.AgentPluginMappingService;
|
import xiaozhi.modules.agent.service.AgentPluginMappingService;
|
||||||
import xiaozhi.modules.agent.service.AgentService;
|
import xiaozhi.modules.agent.service.AgentService;
|
||||||
import xiaozhi.modules.agent.service.AgentTemplateService;
|
import xiaozhi.modules.agent.service.AgentTemplateService;
|
||||||
@@ -39,6 +44,7 @@ public class ConfigServiceImpl implements ConfigService {
|
|||||||
private final RedisUtils redisUtils;
|
private final RedisUtils redisUtils;
|
||||||
private final TimbreService timbreService;
|
private final TimbreService timbreService;
|
||||||
private final AgentPluginMappingService agentPluginMappingService;
|
private final AgentPluginMappingService agentPluginMappingService;
|
||||||
|
private final AgentMcpAccessPointService agentMcpAccessPointService;
|
||||||
|
|
||||||
@Override
|
@Override
|
||||||
public Object getConfig(Boolean isCache) {
|
public Object getConfig(Boolean isCache) {
|
||||||
@@ -144,6 +150,12 @@ public class ConfigServiceImpl implements ConfigService {
|
|||||||
result.put("plugins", pluginParams);
|
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(
|
buildModuleConfig(
|
||||||
|
|||||||
+11
-4
@@ -4,6 +4,7 @@ import java.util.List;
|
|||||||
|
|
||||||
import org.apache.commons.lang3.StringUtils;
|
import org.apache.commons.lang3.StringUtils;
|
||||||
import org.apache.shiro.authz.annotation.RequiresPermissions;
|
import org.apache.shiro.authz.annotation.RequiresPermissions;
|
||||||
|
import org.springframework.beans.BeanUtils;
|
||||||
import org.springframework.web.bind.annotation.GetMapping;
|
import org.springframework.web.bind.annotation.GetMapping;
|
||||||
import org.springframework.web.bind.annotation.PathVariable;
|
import org.springframework.web.bind.annotation.PathVariable;
|
||||||
import org.springframework.web.bind.annotation.PostMapping;
|
import org.springframework.web.bind.annotation.PostMapping;
|
||||||
@@ -14,6 +15,7 @@ import org.springframework.web.bind.annotation.RestController;
|
|||||||
|
|
||||||
import io.swagger.v3.oas.annotations.Operation;
|
import io.swagger.v3.oas.annotations.Operation;
|
||||||
import io.swagger.v3.oas.annotations.tags.Tag;
|
import io.swagger.v3.oas.annotations.tags.Tag;
|
||||||
|
import jakarta.validation.Valid;
|
||||||
import lombok.AllArgsConstructor;
|
import lombok.AllArgsConstructor;
|
||||||
import xiaozhi.common.exception.ErrorCode;
|
import xiaozhi.common.exception.ErrorCode;
|
||||||
import xiaozhi.common.redis.RedisKeys;
|
import xiaozhi.common.redis.RedisKeys;
|
||||||
@@ -22,6 +24,7 @@ import xiaozhi.common.user.UserDetail;
|
|||||||
import xiaozhi.common.utils.Result;
|
import xiaozhi.common.utils.Result;
|
||||||
import xiaozhi.modules.device.dto.DeviceRegisterDTO;
|
import xiaozhi.modules.device.dto.DeviceRegisterDTO;
|
||||||
import xiaozhi.modules.device.dto.DeviceUnBindDTO;
|
import xiaozhi.modules.device.dto.DeviceUnBindDTO;
|
||||||
|
import xiaozhi.modules.device.dto.DeviceUpdateDTO;
|
||||||
import xiaozhi.modules.device.entity.DeviceEntity;
|
import xiaozhi.modules.device.entity.DeviceEntity;
|
||||||
import xiaozhi.modules.device.service.DeviceService;
|
import xiaozhi.modules.device.service.DeviceService;
|
||||||
import xiaozhi.modules.security.user.SecurityUser;
|
import xiaozhi.modules.security.user.SecurityUser;
|
||||||
@@ -80,15 +83,19 @@ public class DeviceController {
|
|||||||
return new Result<Void>();
|
return new Result<Void>();
|
||||||
}
|
}
|
||||||
|
|
||||||
@PutMapping("/enableOta/{id}/{status}")
|
@PutMapping("/update/{id}")
|
||||||
@Operation(summary = "启用/关闭OTA自动升级")
|
@Operation(summary = "更新设备信息")
|
||||||
@RequiresPermissions("sys:role:normal")
|
@RequiresPermissions("sys:role:normal")
|
||||||
public Result<Void> enableOtaUpgrade(@PathVariable String id, @PathVariable Integer status) {
|
public Result<Void> updateDeviceInfo(@PathVariable String id, @Valid @RequestBody DeviceUpdateDTO deviceUpdateDTO) {
|
||||||
DeviceEntity entity = deviceService.selectById(id);
|
DeviceEntity entity = deviceService.selectById(id);
|
||||||
if (entity == null) {
|
if (entity == null) {
|
||||||
return new Result<Void>().error("设备不存在");
|
return new Result<Void>().error("设备不存在");
|
||||||
}
|
}
|
||||||
entity.setAutoUpdate(status);
|
UserDetail user = SecurityUser.getUser();
|
||||||
|
if (!entity.getUserId().equals(user.getId())) {
|
||||||
|
return new Result<Void>().error("设备不存在");
|
||||||
|
}
|
||||||
|
BeanUtils.copyProperties(deviceUpdateDTO, entity);
|
||||||
deviceService.updateById(entity);
|
deviceService.updateById(entity);
|
||||||
return new Result<Void>();
|
return new Result<Void>();
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,29 @@
|
|||||||
|
package xiaozhi.modules.device.dto;
|
||||||
|
|
||||||
|
import jakarta.validation.constraints.Max;
|
||||||
|
import jakarta.validation.constraints.Min;
|
||||||
|
import jakarta.validation.constraints.Size;
|
||||||
|
import lombok.Data;
|
||||||
|
|
||||||
|
import java.io.Serializable;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* 设备更新DTO
|
||||||
|
*/
|
||||||
|
@Data
|
||||||
|
public class DeviceUpdateDTO implements Serializable {
|
||||||
|
/**
|
||||||
|
* 自动更新状态
|
||||||
|
*/
|
||||||
|
@Max(1)
|
||||||
|
@Min(0)
|
||||||
|
private Integer autoUpdate;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* 设备别名
|
||||||
|
*/
|
||||||
|
@Size(max = 64)
|
||||||
|
private String alias;
|
||||||
|
|
||||||
|
private static final long serialVersionUID = 1L;
|
||||||
|
}
|
||||||
+32
@@ -104,6 +104,9 @@ public class SysParamsController {
|
|||||||
// 验证OTA地址
|
// 验证OTA地址
|
||||||
validateOtaUrl(dto.getParamCode(), dto.getParamValue());
|
validateOtaUrl(dto.getParamCode(), dto.getParamValue());
|
||||||
|
|
||||||
|
// 验证MCP地址
|
||||||
|
validateMcpUrl(dto.getParamCode(), dto.getParamValue());
|
||||||
|
|
||||||
sysParamsService.update(dto);
|
sysParamsService.update(dto);
|
||||||
configService.getConfig(false);
|
configService.getConfig(false);
|
||||||
return new Result<Void>();
|
return new Result<Void>();
|
||||||
@@ -195,4 +198,33 @@ public class SysParamsController {
|
|||||||
throw new RenException("OTA接口验证失败:" + e.getMessage());
|
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());
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1 @@
|
|||||||
|
ALTER TABLE ai_agent_plugin_mapping CONVERT TO CHARACTER SET utf8mb4;
|
||||||
@@ -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,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,3 +205,24 @@ databaseChangeLog:
|
|||||||
- sqlFile:
|
- sqlFile:
|
||||||
encoding: utf8
|
encoding: utf8
|
||||||
path: classpath:db/changelog/202506080955.sql
|
path: classpath:db/changelog/202506080955.sql
|
||||||
|
- changeSet:
|
||||||
|
id: 202506161101
|
||||||
|
author: hrz
|
||||||
|
changes:
|
||||||
|
- sqlFile:
|
||||||
|
encoding: utf8
|
||||||
|
path: classpath:db/changelog/202506161101.sql
|
||||||
|
- changeSet:
|
||||||
|
id: 202506191643
|
||||||
|
author: hrz
|
||||||
|
changes:
|
||||||
|
- sqlFile:
|
||||||
|
encoding: utf8
|
||||||
|
path: classpath:db/changelog/202506191643.sql
|
||||||
|
- changeSet:
|
||||||
|
id: 202506261637
|
||||||
|
author: hrz
|
||||||
|
changes:
|
||||||
|
- sqlFile:
|
||||||
|
encoding: utf8
|
||||||
|
path: classpath:db/changelog/202506261637.sql
|
||||||
@@ -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);
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -143,4 +143,34 @@ export default {
|
|||||||
});
|
});
|
||||||
}).send();
|
}).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();
|
||||||
|
},
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -51,10 +51,11 @@ export default {
|
|||||||
});
|
});
|
||||||
}).send();
|
}).send();
|
||||||
},
|
},
|
||||||
enableOtaUpgrade(id, status, callback) {
|
updateDeviceInfo(id, payload, callback) {
|
||||||
RequestService.sendRequest()
|
RequestService.sendRequest()
|
||||||
.url(`${getServiceUrl()}/device/enableOta/${id}/${status}`)
|
.url(`${getServiceUrl()}/device/update/${id}`)
|
||||||
.method('PUT')
|
.method('PUT')
|
||||||
|
.data(payload)
|
||||||
.success((res) => {
|
.success((res) => {
|
||||||
RequestService.clearRequestTime()
|
RequestService.clearRequestTime()
|
||||||
callback(res)
|
callback(res)
|
||||||
@@ -63,7 +64,7 @@ export default {
|
|||||||
console.error('更新OTA状态失败:', err)
|
console.error('更新OTA状态失败:', err)
|
||||||
this.$message.error(err.msg || '更新OTA状态失败')
|
this.$message.error(err.msg || '更新OTA状态失败')
|
||||||
RequestService.reAjaxFun(() => {
|
RequestService.reAjaxFun(() => {
|
||||||
this.enableOtaUpgrade(id, status, callback)
|
this.updateDeviceInfo(id, payload, callback)
|
||||||
})
|
})
|
||||||
}).send()
|
}).send()
|
||||||
},
|
},
|
||||||
|
|||||||
@@ -99,6 +99,60 @@
|
|||||||
</div>
|
</div>
|
||||||
</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">
|
<div class="drawer-footer">
|
||||||
<el-button @click="closeDialog">取消</el-button>
|
<el-button @click="closeDialog">取消</el-button>
|
||||||
<el-button type="primary" @click="saveSelection">保存配置</el-button>
|
<el-button type="primary" @click="saveSelection">保存配置</el-button>
|
||||||
@@ -107,6 +161,8 @@
|
|||||||
</template>
|
</template>
|
||||||
|
|
||||||
<script>
|
<script>
|
||||||
|
import Api from '@/apis/api';
|
||||||
|
|
||||||
export default {
|
export default {
|
||||||
props: {
|
props: {
|
||||||
value: Boolean,
|
value: Boolean,
|
||||||
@@ -117,6 +173,10 @@ export default {
|
|||||||
allFunctions: {
|
allFunctions: {
|
||||||
type: Array,
|
type: Array,
|
||||||
default: () => []
|
default: () => []
|
||||||
|
},
|
||||||
|
agentId: {
|
||||||
|
type: String,
|
||||||
|
required: true
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
data() {
|
data() {
|
||||||
@@ -134,6 +194,10 @@ export default {
|
|||||||
// 添加一个标志位来跟踪是否已经保存
|
// 添加一个标志位来跟踪是否已经保存
|
||||||
hasSaved: false,
|
hasSaved: false,
|
||||||
loading: false,
|
loading: false,
|
||||||
|
|
||||||
|
mcpUrl: "",
|
||||||
|
mcpStatus: "disconnected",
|
||||||
|
mcpTools: [],
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
computed: {
|
computed: {
|
||||||
@@ -177,6 +241,10 @@ export default {
|
|||||||
});
|
});
|
||||||
// 右侧默认指向第一个
|
// 右侧默认指向第一个
|
||||||
this.currentFunction = this.selectedList[0] || null;
|
this.currentFunction = this.selectedList[0] || null;
|
||||||
|
|
||||||
|
// 加载MCP数据
|
||||||
|
this.loadMcpAddress();
|
||||||
|
this.loadMcpTools();
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
dialogVisible(newVal) {
|
dialogVisible(newVal) {
|
||||||
@@ -184,6 +252,60 @@ export default {
|
|||||||
}
|
}
|
||||||
},
|
},
|
||||||
methods: {
|
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) {
|
flushArray(key) {
|
||||||
const text = this.textCache[key] || '';
|
const text = this.textCache[key] || '';
|
||||||
const arr = text
|
const arr = text
|
||||||
@@ -298,7 +420,7 @@ export default {
|
|||||||
display: grid;
|
display: grid;
|
||||||
grid-template-columns: max-content max-content 1fr;
|
grid-template-columns: max-content max-content 1fr;
|
||||||
gap: 12px;
|
gap: 12px;
|
||||||
height: calc(70vh - 60px);
|
height: calc(58vh);
|
||||||
}
|
}
|
||||||
|
|
||||||
.custom-header {
|
.custom-header {
|
||||||
@@ -514,4 +636,202 @@ export default {
|
|||||||
::v-deep .el-checkbox__label {
|
::v-deep .el-checkbox__label {
|
||||||
display: none;
|
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>
|
</style>
|
||||||
@@ -1,12 +1,12 @@
|
|||||||
<template>
|
<template>
|
||||||
<div class="welcome">
|
<div class="welcome">
|
||||||
<HeaderBar />
|
<HeaderBar/>
|
||||||
|
|
||||||
<div class="operation-bar">
|
<div class="operation-bar">
|
||||||
<h2 class="page-title">设备管理</h2>
|
<h2 class="page-title">设备管理</h2>
|
||||||
<div class="right-operations">
|
<div class="right-operations">
|
||||||
<el-input placeholder="请输入设备型号或Mac地址查询" v-model="searchKeyword" class="search-input"
|
<el-input placeholder="请输入设备型号或Mac地址查询" v-model="searchKeyword" class="search-input"
|
||||||
@keyup.enter.native="handleSearch" clearable />
|
@keyup.enter.native="handleSearch" clearable/>
|
||||||
<el-button class="btn-search" @click="handleSearch">搜索</el-button>
|
<el-button class="btn-search" @click="handleSearch">搜索</el-button>
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
@@ -16,8 +16,9 @@
|
|||||||
<div class="content-area">
|
<div class="content-area">
|
||||||
<el-card class="device-card" shadow="never">
|
<el-card class="device-card" shadow="never">
|
||||||
<el-table ref="deviceTable" :data="paginatedDeviceList" class="transparent-table"
|
<el-table ref="deviceTable" :data="paginatedDeviceList" class="transparent-table"
|
||||||
:header-cell-class-name="headerCellClassName" v-loading="loading" element-loading-text="拼命加载中"
|
:header-cell-class-name="headerCellClassName" v-loading="loading"
|
||||||
element-loading-spinner="el-icon-loading" element-loading-background="rgba(255, 255, 255, 0.7)">
|
element-loading-text="拼命加载中"
|
||||||
|
element-loading-spinner="el-icon-loading" element-loading-background="rgba(255, 255, 255, 0.7)">
|
||||||
<el-table-column label="选择" align="center" width="120">
|
<el-table-column label="选择" align="center" width="120">
|
||||||
<template slot-scope="scope">
|
<template slot-scope="scope">
|
||||||
<el-checkbox v-model="scope.row.selected"></el-checkbox>
|
<el-checkbox v-model="scope.row.selected"></el-checkbox>
|
||||||
@@ -33,22 +34,32 @@
|
|||||||
<el-table-column label="绑定时间" prop="bindTime" align="center"></el-table-column>
|
<el-table-column label="绑定时间" prop="bindTime" align="center"></el-table-column>
|
||||||
<el-table-column label="最近对话" prop="lastConversation" align="center"></el-table-column>
|
<el-table-column label="最近对话" prop="lastConversation" align="center"></el-table-column>
|
||||||
<el-table-column label="备注" align="center">
|
<el-table-column label="备注" align="center">
|
||||||
<template slot-scope="scope">
|
<template #default="{ row }">
|
||||||
<el-input v-if="scope.row.isEdit" v-model="scope.row.remark" size="mini"
|
<el-input
|
||||||
@blur="stopEditRemark(scope.$index)"></el-input>
|
v-show="row.isEdit"
|
||||||
<span v-else>
|
v-model="row.remark"
|
||||||
<i v-if="!scope.row.remark" class="el-icon-edit"
|
size="mini"
|
||||||
@click="startEditRemark(scope.$index, scope.row)"></i>
|
maxlength="64"
|
||||||
<span v-else @click="startEditRemark(scope.$index, scope.row)">
|
show-word-limit
|
||||||
{{ scope.row.remark }}
|
@blur="onRemarkBlur(row)"
|
||||||
</span>
|
@keyup.enter.native="onRemarkEnter(row)"
|
||||||
|
/>
|
||||||
|
<span v-show="!row.isEdit" class="remark-view">
|
||||||
|
<i
|
||||||
|
class="el-icon-edit"
|
||||||
|
@click="row.isEdit = true"
|
||||||
|
style="cursor: pointer;"
|
||||||
|
></i>
|
||||||
|
<span @click="row.isEdit = true">
|
||||||
|
{{ row.remark || '—' }}
|
||||||
</span>
|
</span>
|
||||||
|
</span>
|
||||||
</template>
|
</template>
|
||||||
</el-table-column>
|
</el-table-column>
|
||||||
<el-table-column label="OTA升级" align="center">
|
<el-table-column label="OTA升级" align="center">
|
||||||
<template slot-scope="scope">
|
<template slot-scope="scope">
|
||||||
<el-switch v-model="scope.row.otaSwitch" size="mini" active-color="#13ce66" inactive-color="#ff4949"
|
<el-switch v-model="scope.row.otaSwitch" size="mini" active-color="#13ce66" inactive-color="#ff4949"
|
||||||
@change="handleOtaSwitchChange(scope.row)"></el-switch>
|
@change="handleOtaSwitchChange(scope.row)"></el-switch>
|
||||||
</template>
|
</template>
|
||||||
</el-table-column>
|
</el-table-column>
|
||||||
<el-table-column label="操作" align="center">
|
<el-table-column label="操作" align="center">
|
||||||
@@ -78,7 +89,7 @@
|
|||||||
<button class="pagination-btn" :disabled="currentPage === 1" @click="goFirst">首页</button>
|
<button class="pagination-btn" :disabled="currentPage === 1" @click="goFirst">首页</button>
|
||||||
<button class="pagination-btn" :disabled="currentPage === 1" @click="goPrev">上一页</button>
|
<button class="pagination-btn" :disabled="currentPage === 1" @click="goPrev">上一页</button>
|
||||||
<button v-for="page in visiblePages" :key="page" class="pagination-btn"
|
<button v-for="page in visiblePages" :key="page" class="pagination-btn"
|
||||||
:class="{ active: page === currentPage }" @click="goToPage(page)">
|
:class="{ active: page === currentPage }" @click="goToPage(page)">
|
||||||
{{ page }}
|
{{ page }}
|
||||||
</button>
|
</button>
|
||||||
<button class="pagination-btn" :disabled="currentPage === pageCount" @click="goNext">下一页</button>
|
<button class="pagination-btn" :disabled="currentPage === pageCount" @click="goNext">下一页</button>
|
||||||
@@ -91,7 +102,7 @@
|
|||||||
</div>
|
</div>
|
||||||
|
|
||||||
<AddDeviceDialog :visible.sync="addDeviceDialogVisible" :agent-id="currentAgentId"
|
<AddDeviceDialog :visible.sync="addDeviceDialogVisible" :agent-id="currentAgentId"
|
||||||
@refresh="fetchBindDevices(currentAgentId)" />
|
@refresh="fetchBindDevices(currentAgentId)"/>
|
||||||
|
|
||||||
</div>
|
</div>
|
||||||
</template>
|
</template>
|
||||||
@@ -102,7 +113,7 @@ import AddDeviceDialog from "@/components/AddDeviceDialog.vue";
|
|||||||
import HeaderBar from "@/components/HeaderBar.vue";
|
import HeaderBar from "@/components/HeaderBar.vue";
|
||||||
|
|
||||||
export default {
|
export default {
|
||||||
components: { HeaderBar, AddDeviceDialog },
|
components: {HeaderBar, AddDeviceDialog},
|
||||||
data() {
|
data() {
|
||||||
return {
|
return {
|
||||||
addDeviceDialogVisible: false,
|
addDeviceDialogVisible: false,
|
||||||
@@ -125,18 +136,15 @@ export default {
|
|||||||
const keyword = this.activeSearchKeyword.toLowerCase();
|
const keyword = this.activeSearchKeyword.toLowerCase();
|
||||||
if (!keyword) return this.deviceList;
|
if (!keyword) return this.deviceList;
|
||||||
return this.deviceList.filter(device =>
|
return this.deviceList.filter(device =>
|
||||||
(device.model && device.model.toLowerCase().includes(keyword)) ||
|
(device.model && device.model.toLowerCase().includes(keyword)) ||
|
||||||
(device.macAddress && device.macAddress.toLowerCase().includes(keyword))
|
(device.macAddress && device.macAddress.toLowerCase().includes(keyword))
|
||||||
);
|
);
|
||||||
},
|
},
|
||||||
|
|
||||||
paginatedDeviceList() {
|
paginatedDeviceList() {
|
||||||
const start = (this.currentPage - 1) * this.pageSize;
|
const start = (this.currentPage - 1) * this.pageSize;
|
||||||
const end = start + this.pageSize;
|
const end = start + this.pageSize;
|
||||||
return this.filteredDeviceList.slice(start, end).map(item => ({
|
return this.filteredDeviceList.slice(start, end);
|
||||||
...item,
|
|
||||||
selected: false
|
|
||||||
}));
|
|
||||||
},
|
},
|
||||||
pageCount() {
|
pageCount() {
|
||||||
return Math.ceil(this.filteredDeviceList.length / this.pageSize);
|
return Math.ceil(this.filteredDeviceList.length / this.pageSize);
|
||||||
@@ -212,11 +220,10 @@ export default {
|
|||||||
this.batchUnbindDevices(deviceIds);
|
this.batchUnbindDevices(deviceIds);
|
||||||
});
|
});
|
||||||
},
|
},
|
||||||
|
|
||||||
batchUnbindDevices(deviceIds) {
|
batchUnbindDevices(deviceIds) {
|
||||||
const promises = deviceIds.map(id => {
|
const promises = deviceIds.map(id => {
|
||||||
return new Promise((resolve, reject) => {
|
return new Promise((resolve, reject) => {
|
||||||
Api.device.unbindDevice(id, ({ data }) => {
|
Api.device.unbindDevice(id, ({data}) => {
|
||||||
if (data.code === 0) {
|
if (data.code === 0) {
|
||||||
resolve();
|
resolve();
|
||||||
} else {
|
} else {
|
||||||
@@ -225,33 +232,61 @@ export default {
|
|||||||
});
|
});
|
||||||
});
|
});
|
||||||
});
|
});
|
||||||
|
|
||||||
Promise.all(promises)
|
Promise.all(promises)
|
||||||
.then(() => {
|
.then(() => {
|
||||||
this.$message.success({
|
this.$message.success({
|
||||||
message: `成功解绑 ${deviceIds.length} 台设备`,
|
message: `成功解绑 ${deviceIds.length} 台设备`,
|
||||||
showClose: true
|
showClose: true
|
||||||
|
});
|
||||||
|
this.fetchBindDevices(this.currentAgentId);
|
||||||
|
this.selectedDevices = [];
|
||||||
|
this.isAllSelected = false;
|
||||||
|
})
|
||||||
|
.catch(error => {
|
||||||
|
this.$message.error({
|
||||||
|
message: error || '批量解绑过程中出现错误',
|
||||||
|
showClose: true
|
||||||
|
});
|
||||||
});
|
});
|
||||||
this.fetchBindDevices(this.currentAgentId);
|
|
||||||
this.selectedDevices = [];
|
|
||||||
this.isAllSelected = false;
|
|
||||||
})
|
|
||||||
.catch(error => {
|
|
||||||
this.$message.error({
|
|
||||||
message: error || '批量解绑过程中出现错误',
|
|
||||||
showClose: true
|
|
||||||
});
|
|
||||||
});
|
|
||||||
},
|
},
|
||||||
|
|
||||||
handleAddDevice() {
|
handleAddDevice() {
|
||||||
this.addDeviceDialogVisible = true;
|
this.addDeviceDialogVisible = true;
|
||||||
},
|
},
|
||||||
startEditRemark(index, row) {
|
submitRemark(row) {
|
||||||
this.deviceList[index].isEdit = true;
|
if (row._submitting) return;
|
||||||
|
|
||||||
|
const text = (row.remark || '').trim();
|
||||||
|
if (text.length > 64) {
|
||||||
|
this.$message.warning('备注不能超过 64 字符');
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
if (text === row._originalRemark) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
row._submitting = true;
|
||||||
|
this.updateDeviceInfo(row.device_id, { alias: text }, (ok, resp) => {
|
||||||
|
if (ok) {
|
||||||
|
row._originalRemark = text;
|
||||||
|
this.$message.success('备注已保存');
|
||||||
|
} else {
|
||||||
|
row.remark = row._originalRemark;
|
||||||
|
this.$message.error(resp.msg || '备注保存失败');
|
||||||
|
}
|
||||||
|
row._submitting = false;
|
||||||
|
});
|
||||||
},
|
},
|
||||||
stopEditRemark(index) {
|
// 备注输入框:失焦时提交
|
||||||
this.deviceList[index].isEdit = false;
|
onRemarkBlur(row) {
|
||||||
|
row.isEdit = false;
|
||||||
|
setTimeout(() => {
|
||||||
|
this.submitRemark(row);
|
||||||
|
}, 100); // 延迟 100ms,避开 enter+blur 同时触发的窗口
|
||||||
|
},
|
||||||
|
// 备注输入框:按回车时提交
|
||||||
|
onRemarkEnter(row) {
|
||||||
|
row.isEdit = false;
|
||||||
|
this.submitRemark(row);
|
||||||
},
|
},
|
||||||
handleUnbind(device_id) {
|
handleUnbind(device_id) {
|
||||||
this.$confirm('确认要解绑该设备吗?', '警告', {
|
this.$confirm('确认要解绑该设备吗?', '警告', {
|
||||||
@@ -259,7 +294,7 @@ export default {
|
|||||||
cancelButtonText: '取消',
|
cancelButtonText: '取消',
|
||||||
type: 'warning'
|
type: 'warning'
|
||||||
}).then(() => {
|
}).then(() => {
|
||||||
Api.device.unbindDevice(device_id, ({ data }) => {
|
Api.device.unbindDevice(device_id, ({data}) => {
|
||||||
if (data.code === 0) {
|
if (data.code === 0) {
|
||||||
this.$message.success({
|
this.$message.success({
|
||||||
message: '设备解绑成功',
|
message: '设备解绑成功',
|
||||||
@@ -290,7 +325,7 @@ export default {
|
|||||||
|
|
||||||
fetchBindDevices(agentId) {
|
fetchBindDevices(agentId) {
|
||||||
this.loading = true;
|
this.loading = true;
|
||||||
Api.device.getAgentBindDevices(agentId, ({ data }) => {
|
Api.device.getAgentBindDevices(agentId, ({data}) => {
|
||||||
this.loading = false;
|
this.loading = false;
|
||||||
if (data.code === 0) {
|
if (data.code === 0) {
|
||||||
this.deviceList = data.data.map(device => {
|
this.deviceList = data.data.map(device => {
|
||||||
@@ -302,12 +337,14 @@ export default {
|
|||||||
bindTime: device.createDate,
|
bindTime: device.createDate,
|
||||||
lastConversation: device.lastConnectedAt,
|
lastConversation: device.lastConnectedAt,
|
||||||
remark: device.alias,
|
remark: device.alias,
|
||||||
|
_originalRemark: device.alias,
|
||||||
isEdit: false,
|
isEdit: false,
|
||||||
|
_submitting: false,
|
||||||
otaSwitch: device.autoUpdate === 1,
|
otaSwitch: device.autoUpdate === 1,
|
||||||
rawBindTime: new Date(device.createDate).getTime()
|
rawBindTime: new Date(device.createDate).getTime()
|
||||||
};
|
};
|
||||||
})
|
})
|
||||||
.sort((a, b) => a.rawBindTime - b.rawBindTime);
|
.sort((a, b) => a.rawBindTime - b.rawBindTime);
|
||||||
this.activeSearchKeyword = "";
|
this.activeSearchKeyword = "";
|
||||||
this.searchKeyword = "";
|
this.searchKeyword = "";
|
||||||
} else {
|
} else {
|
||||||
@@ -315,7 +352,7 @@ export default {
|
|||||||
}
|
}
|
||||||
});
|
});
|
||||||
},
|
},
|
||||||
headerCellClassName({ columnIndex }) {
|
headerCellClassName({columnIndex}) {
|
||||||
if (columnIndex === 0) {
|
if (columnIndex === 0) {
|
||||||
return "custom-selection-header";
|
return "custom-selection-header";
|
||||||
}
|
}
|
||||||
@@ -325,14 +362,19 @@ export default {
|
|||||||
const firmwareType = this.firmwareTypes.find(item => item.key === type)
|
const firmwareType = this.firmwareTypes.find(item => item.key === type)
|
||||||
return firmwareType ? firmwareType.name : type
|
return firmwareType ? firmwareType.name : type
|
||||||
},
|
},
|
||||||
|
updateDeviceInfo(device_id, payload, callback) {
|
||||||
|
return Api.device.updateDeviceInfo(device_id, payload, ({data}) => {
|
||||||
|
callback(data.code === 0, data);
|
||||||
|
})
|
||||||
|
},
|
||||||
handleOtaSwitchChange(row) {
|
handleOtaSwitchChange(row) {
|
||||||
Api.device.enableOtaUpgrade(row.device_id, row.otaSwitch ? 1 : 0, ({ data }) => {
|
this.updateDeviceInfo(row.device_id, {autoUpdate: row.otaSwitch ? 1 : 0}, (result, {msg}) => {
|
||||||
if (data.code === 0) {
|
if (result) {
|
||||||
this.$message.success(row.otaSwitch ? '已设置成自动升级' : '已关闭自动升级')
|
this.$message.success(row.otaSwitch ? '已设置成自动升级' : '已关闭自动升级');
|
||||||
} else {
|
return;
|
||||||
row.otaSwitch = !row.otaSwitch
|
|
||||||
this.$message.error(data.msg || '操作失败')
|
|
||||||
}
|
}
|
||||||
|
row.otaSwitch = !row.otaSwitch
|
||||||
|
this.$message.error(msg || '操作失败')
|
||||||
})
|
})
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
@@ -645,7 +687,6 @@ export default {
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
:deep(.el-table .el-button--text) {
|
:deep(.el-table .el-button--text) {
|
||||||
color: #7079aa;
|
color: #7079aa;
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -132,7 +132,7 @@
|
|||||||
</div>
|
</div>
|
||||||
|
|
||||||
<function-dialog v-model="showFunctionDialog" :functions="currentFunctions" :all-functions="allFunctions"
|
<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>
|
</div>
|
||||||
</template>
|
</template>
|
||||||
|
|
||||||
|
|||||||
@@ -104,7 +104,8 @@ wakeup_words:
|
|||||||
- "喵喵同学"
|
- "喵喵同学"
|
||||||
- "小滨小滨"
|
- "小滨小滨"
|
||||||
- "小冰小冰"
|
- "小冰小冰"
|
||||||
|
# MCP接入点地址
|
||||||
|
mcp_endpoint: 你的接入点 websocket地址
|
||||||
# 插件的基础配置
|
# 插件的基础配置
|
||||||
plugins:
|
plugins:
|
||||||
# 获取天气插件的配置,这里填写你的api_key
|
# 获取天气插件的配置,这里填写你的api_key
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ from config.config_loader import load_config
|
|||||||
from config.settings import check_config_file
|
from config.settings import check_config_file
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
|
|
||||||
SERVER_VERSION = "0.5.6"
|
SERVER_VERSION = "0.6.1"
|
||||||
_logger_initialized = False
|
_logger_initialized = False
|
||||||
|
|
||||||
|
|
||||||
@@ -88,7 +88,7 @@ def setup_logging():
|
|||||||
# 输出到文件 - 统一目录,按大小轮转
|
# 输出到文件 - 统一目录,按大小轮转
|
||||||
# 日志文件完整路径
|
# 日志文件完整路径
|
||||||
log_file_path = os.path.join(log_dir, log_file)
|
log_file_path = os.path.join(log_dir, log_file)
|
||||||
|
|
||||||
# 添加日志处理器
|
# 添加日志处理器
|
||||||
logger.add(
|
logger.add(
|
||||||
log_file_path,
|
log_file_path,
|
||||||
@@ -146,14 +146,14 @@ def update_module_string(selected_module_str):
|
|||||||
level=log_config.get("log_level", "INFO"),
|
level=log_config.get("log_level", "INFO"),
|
||||||
filter=formatter,
|
filter=formatter,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 更新文件日志配置 - 统一目录,按大小轮转
|
# 更新文件日志配置 - 统一目录,按大小轮转
|
||||||
log_dir = log_config.get("log_dir", "tmp")
|
log_dir = log_config.get("log_dir", "tmp")
|
||||||
log_file = log_config.get("log_file", "server.log")
|
log_file = log_config.get("log_file", "server.log")
|
||||||
|
|
||||||
# 日志文件完整路径
|
# 日志文件完整路径
|
||||||
log_file_path = os.path.join(log_dir, log_file)
|
log_file_path = os.path.join(log_dir, log_file)
|
||||||
|
|
||||||
logger.add(
|
logger.add(
|
||||||
log_file_path,
|
log_file_path,
|
||||||
format=log_format_file,
|
format=log_format_file,
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ from config.config_loader import get_private_config_from_api
|
|||||||
from core.utils.auth import AuthToken
|
from core.utils.auth import AuthToken
|
||||||
import base64
|
import base64
|
||||||
from typing import Tuple, Optional
|
from typing import Tuple, Optional
|
||||||
|
from plugins_func.register import Action
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
|
|
||||||
@@ -122,7 +123,8 @@ class VisionHandler:
|
|||||||
|
|
||||||
return_json = {
|
return_json = {
|
||||||
"success": True,
|
"success": True,
|
||||||
"result": result,
|
"action": Action.RESPONSE.name,
|
||||||
|
"response": result,
|
||||||
}
|
}
|
||||||
|
|
||||||
response = web.Response(
|
response = web.Response(
|
||||||
|
|||||||
@@ -10,7 +10,6 @@ import threading
|
|||||||
import traceback
|
import traceback
|
||||||
import subprocess
|
import subprocess
|
||||||
import websockets
|
import websockets
|
||||||
from core.handle.mcpHandle import call_mcp_tool
|
|
||||||
from core.utils.util import (
|
from core.utils.util import (
|
||||||
extract_json_from_string,
|
extract_json_from_string,
|
||||||
check_vad_update,
|
check_vad_update,
|
||||||
@@ -18,7 +17,6 @@ from core.utils.util import (
|
|||||||
filter_sensitive_info,
|
filter_sensitive_info,
|
||||||
)
|
)
|
||||||
from typing import Dict, Any
|
from typing import Dict, Any
|
||||||
from core.mcp.manager import MCPManager
|
|
||||||
from core.utils.modules_initialize import (
|
from core.utils.modules_initialize import (
|
||||||
initialize_modules,
|
initialize_modules,
|
||||||
initialize_tts,
|
initialize_tts,
|
||||||
@@ -30,7 +28,7 @@ from concurrent.futures import ThreadPoolExecutor
|
|||||||
from core.utils.dialogue import Message, Dialogue
|
from core.utils.dialogue import Message, Dialogue
|
||||||
from core.providers.asr.dto.dto import InterfaceType
|
from core.providers.asr.dto.dto import InterfaceType
|
||||||
from core.handle.textHandle import handleTextMessage
|
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.loadplugins import auto_import_modules
|
||||||
from plugins_func.register import Action, ActionResponse
|
from plugins_func.register import Action, ActionResponse
|
||||||
from core.auth import AuthMiddleware, AuthenticationError
|
from core.auth import AuthMiddleware, AuthenticationError
|
||||||
@@ -112,8 +110,7 @@ class ConnectionHandler:
|
|||||||
# vad相关变量
|
# vad相关变量
|
||||||
self.client_audio_buffer = bytearray()
|
self.client_audio_buffer = bytearray()
|
||||||
self.client_have_voice = False
|
self.client_have_voice = False
|
||||||
self.client_have_voice_last_time = 0.0
|
self.last_activity_time = 0.0 # 统一的活动时间戳(毫秒)
|
||||||
self.client_no_voice_last_time = 0.0
|
|
||||||
self.client_voice_stop = False
|
self.client_voice_stop = False
|
||||||
|
|
||||||
# asr相关变量
|
# asr相关变量
|
||||||
@@ -144,10 +141,10 @@ class ConnectionHandler:
|
|||||||
self.load_function_plugin = False
|
self.load_function_plugin = False
|
||||||
self.intent_type = "nointent"
|
self.intent_type = "nointent"
|
||||||
|
|
||||||
self.timeout_task = None
|
|
||||||
self.timeout_seconds = (
|
self.timeout_seconds = (
|
||||||
int(self.config.get("close_connection_no_voice_time", 120)) + 60
|
int(self.config.get("close_connection_no_voice_time", 120)) + 60
|
||||||
) # 在原来第一道关闭的基础上加60秒,进行二道关闭
|
) # 在原来第一道关闭的基础上加60秒,进行二道关闭
|
||||||
|
self.timeout_task = None
|
||||||
|
|
||||||
# {"mcp":true} 表示启用MCP功能
|
# {"mcp":true} 表示启用MCP功能
|
||||||
self.features = None
|
self.features = None
|
||||||
@@ -188,6 +185,9 @@ class ConnectionHandler:
|
|||||||
self.websocket = ws
|
self.websocket = ws
|
||||||
self.device_id = self.headers.get("device-id", None)
|
self.device_id = self.headers.get("device-id", None)
|
||||||
|
|
||||||
|
# 初始化活动时间戳
|
||||||
|
self.last_activity_time = time.time() * 1000
|
||||||
|
|
||||||
# 启动超时检查任务
|
# 启动超时检查任务
|
||||||
self.timeout_task = asyncio.create_task(self._check_timeout())
|
self.timeout_task = asyncio.create_task(self._check_timeout())
|
||||||
|
|
||||||
@@ -242,18 +242,10 @@ class ConnectionHandler:
|
|||||||
# 立即关闭连接,不等待记忆保存完成
|
# 立即关闭连接,不等待记忆保存完成
|
||||||
await self.close(ws)
|
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):
|
async def _route_message(self, message):
|
||||||
"""消息路由"""
|
"""消息路由"""
|
||||||
# 重置超时计时器
|
|
||||||
await self.reset_timeout()
|
|
||||||
|
|
||||||
if isinstance(message, str):
|
if isinstance(message, str):
|
||||||
|
self.last_activity_time = time.time() * 1000
|
||||||
await handleTextMessage(self, message)
|
await handleTextMessage(self, message)
|
||||||
elif isinstance(message, bytes):
|
elif isinstance(message, bytes):
|
||||||
if self.vad is None:
|
if self.vad is None:
|
||||||
@@ -474,6 +466,8 @@ class ConnectionHandler:
|
|||||||
self.max_output_size = int(private_config["device_max_output_size"])
|
self.max_output_size = int(private_config["device_max_output_size"])
|
||||||
if private_config.get("chat_history_conf", None) is not None:
|
if private_config.get("chat_history_conf", None) is not None:
|
||||||
self.chat_history_conf = int(private_config["chat_history_conf"])
|
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:
|
try:
|
||||||
modules = initialize_modules(
|
modules = initialize_modules(
|
||||||
self.logger,
|
self.logger,
|
||||||
@@ -585,14 +579,12 @@ class ConnectionHandler:
|
|||||||
self.intent.set_llm(self.llm)
|
self.intent.set_llm(self.llm)
|
||||||
self.logger.bind(tag=TAG).info("使用主LLM作为意图识别模型")
|
self.logger.bind(tag=TAG).info("使用主LLM作为意图识别模型")
|
||||||
|
|
||||||
"""加载插件"""
|
"""加载统一工具处理器"""
|
||||||
self.func_handler = FunctionHandler(self)
|
self.func_handler = UnifiedToolHandler(self)
|
||||||
self.mcp_manager = MCPManager(self)
|
|
||||||
|
|
||||||
"""加载MCP工具"""
|
# 异步初始化工具处理器
|
||||||
asyncio.run_coroutine_threadsafe(
|
if hasattr(self, "loop") and self.loop:
|
||||||
self.mcp_manager.initialize_servers(), self.loop
|
asyncio.run_coroutine_threadsafe(self.func_handler._initialize(), self.loop)
|
||||||
)
|
|
||||||
|
|
||||||
def change_system_prompt(self, prompt):
|
def change_system_prompt(self, prompt):
|
||||||
self.prompt = prompt
|
self.prompt = prompt
|
||||||
@@ -610,12 +602,6 @@ class ConnectionHandler:
|
|||||||
functions = None
|
functions = None
|
||||||
if self.intent_type == "function_call" and hasattr(self, "func_handler"):
|
if self.intent_type == "function_call" and hasattr(self, "func_handler"):
|
||||||
functions = self.func_handler.get_functions()
|
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 = []
|
response_message = []
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -627,8 +613,7 @@ class ConnectionHandler:
|
|||||||
)
|
)
|
||||||
memory_str = future.result()
|
memory_str = future.result()
|
||||||
|
|
||||||
uuid_str = str(uuid.uuid4()).replace("-", "")
|
self.sentence_id = str(uuid.uuid4().hex)
|
||||||
self.sentence_id = uuid_str
|
|
||||||
|
|
||||||
if self.intent_type == "function_call" and functions is not None:
|
if self.intent_type == "function_call" and functions is not None:
|
||||||
# 使用支持functions的streaming接口
|
# 使用支持functions的streaming接口
|
||||||
@@ -733,37 +718,13 @@ class ConnectionHandler:
|
|||||||
"arguments": function_arguments,
|
"arguments": function_arguments,
|
||||||
}
|
}
|
||||||
|
|
||||||
# 处理Server端MCP工具调用
|
# 使用统一工具处理器处理所有工具调用
|
||||||
if self.mcp_manager.is_mcp_tool(function_name):
|
result = asyncio.run_coroutine_threadsafe(
|
||||||
result = self._handle_mcp_tool_call(function_call_data)
|
self.func_handler.handle_llm_function_call(
|
||||||
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(
|
|
||||||
self, function_call_data
|
self, function_call_data
|
||||||
)
|
),
|
||||||
|
self.loop,
|
||||||
|
).result()
|
||||||
self._handle_function_result(result, function_call_data)
|
self._handle_function_result(result, function_call_data)
|
||||||
|
|
||||||
# 存储对话内容
|
# 存储对话内容
|
||||||
@@ -786,48 +747,6 @@ class ConnectionHandler:
|
|||||||
|
|
||||||
return True
|
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):
|
def _handle_function_result(self, result, function_call_data):
|
||||||
if result.action == Action.RESPONSE: # 直接回复前端
|
if result.action == Action.RESPONSE: # 直接回复前端
|
||||||
text = result.response
|
text = result.response
|
||||||
@@ -867,7 +786,7 @@ class ConnectionHandler:
|
|||||||
)
|
)
|
||||||
self.chat(text, tool_call=True)
|
self.chat(text, tool_call=True)
|
||||||
elif result.action == Action.NOTFOUND or result.action == Action.ERROR:
|
elif result.action == Action.NOTFOUND or result.action == Action.ERROR:
|
||||||
text = result.result
|
text = result.response if result.response else result.result
|
||||||
self.tts.tts_one_sentence(self, ContentType.TEXT, content_detail=text)
|
self.tts.tts_one_sentence(self, ContentType.TEXT, content_detail=text)
|
||||||
self.dialogue.put(Message(role="assistant", content=text))
|
self.dialogue.put(Message(role="assistant", content=text))
|
||||||
else:
|
else:
|
||||||
@@ -922,9 +841,9 @@ class ConnectionHandler:
|
|||||||
self.timeout_task.cancel()
|
self.timeout_task.cancel()
|
||||||
self.timeout_task = None
|
self.timeout_task = None
|
||||||
|
|
||||||
# 清理MCP资源
|
# 清理工具处理器资源
|
||||||
if hasattr(self, "mcp_manager") and self.mcp_manager:
|
if hasattr(self, "func_handler") and self.func_handler:
|
||||||
await self.mcp_manager.cleanup_all()
|
await self.func_handler.cleanup()
|
||||||
|
|
||||||
# 触发停止事件
|
# 触发停止事件
|
||||||
if self.stop_event:
|
if self.stop_event:
|
||||||
@@ -976,7 +895,6 @@ class ConnectionHandler:
|
|||||||
def reset_vad_states(self):
|
def reset_vad_states(self):
|
||||||
self.client_audio_buffer = bytearray()
|
self.client_audio_buffer = bytearray()
|
||||||
self.client_have_voice = False
|
self.client_have_voice = False
|
||||||
self.client_have_voice_last_time = 0
|
|
||||||
self.client_voice_stop = False
|
self.client_voice_stop = False
|
||||||
self.logger.bind(tag=TAG).debug("VAD states reset.")
|
self.logger.bind(tag=TAG).debug("VAD states reset.")
|
||||||
|
|
||||||
@@ -995,10 +913,18 @@ class ConnectionHandler:
|
|||||||
"""检查连接超时"""
|
"""检查连接超时"""
|
||||||
try:
|
try:
|
||||||
while not self.stop_event.is_set():
|
while not self.stop_event.is_set():
|
||||||
await asyncio.sleep(self.timeout_seconds)
|
# 检查是否超时(只有在时间戳已初始化的情况下)
|
||||||
if not self.stop_event.is_set():
|
if self.last_activity_time > 0.0:
|
||||||
self.logger.bind(tag=TAG).info("连接超时,准备关闭")
|
current_time = time.time() * 1000
|
||||||
await self.close(self.websocket)
|
if (
|
||||||
break
|
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:
|
except Exception as e:
|
||||||
self.logger.bind(tag=TAG).error(f"超时检查任务出错: {e}")
|
self.logger.bind(tag=TAG).error(f"超时检查任务出错: {e}")
|
||||||
|
|||||||
@@ -1,89 +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 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.handle.sendAudioHandle import sendAudioMessage, send_stt_message
|
||||||
from core.utils.util import remove_punctuation_and_length, opus_datas_to_wav_bytes
|
from core.utils.util import remove_punctuation_and_length, opus_datas_to_wav_bytes
|
||||||
from core.providers.tts.dto.dto import ContentType, SentenceType
|
from core.providers.tts.dto.dto import ContentType, SentenceType
|
||||||
from core.handle.mcpHandle import (
|
from core.providers.tools.device_mcp import (
|
||||||
MCPClient,
|
MCPClient,
|
||||||
send_mcp_initialize_message,
|
send_mcp_initialize_message,
|
||||||
send_mcp_tools_list_request,
|
send_mcp_tools_list_request,
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ from core.handle.helloHandle import checkWakeupWords
|
|||||||
from core.utils.util import remove_punctuation_and_length
|
from core.utils.util import remove_punctuation_and_length
|
||||||
from core.providers.tts.dto.dto import ContentType
|
from core.providers.tts.dto.dto import ContentType
|
||||||
from core.utils.dialogue import Message
|
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 plugins_func.register import Action, ActionResponse
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
@@ -106,36 +106,18 @@ async def process_intent_result(conn, intent_result, original_text):
|
|||||||
def process_function_call():
|
def process_function_call():
|
||||||
conn.dialogue.put(Message(role="user", content=original_text))
|
conn.dialogue.put(Message(role="user", content=original_text))
|
||||||
|
|
||||||
# 处理Server端MCP工具调用
|
# 使用统一工具处理器处理所有工具调用
|
||||||
if conn.mcp_manager.is_mcp_tool(function_name):
|
try:
|
||||||
result = conn._handle_mcp_tool_call(function_call_data)
|
result = asyncio.run_coroutine_threadsafe(
|
||||||
elif hasattr(conn, "mcp_client") and conn.mcp_client.has_tool(
|
conn.func_handler.handle_llm_function_call(
|
||||||
function_name
|
conn, function_call_data
|
||||||
):
|
),
|
||||||
# 如果是小智端MCP工具调用
|
conn.loop,
|
||||||
conn.logger.bind(tag=TAG).debug(
|
).result()
|
||||||
f"调用小智端MCP工具: {function_name}, 参数: {function_args}"
|
except Exception as e:
|
||||||
)
|
conn.logger.bind(tag=TAG).error(f"工具调用失败: {e}")
|
||||||
try:
|
result = ActionResponse(
|
||||||
result = asyncio.run_coroutine_threadsafe(
|
action=Action.ERROR, result=str(e), response=str(e)
|
||||||
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
|
|
||||||
)
|
)
|
||||||
|
|
||||||
if result:
|
if result:
|
||||||
|
|||||||
@@ -1,421 +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"):
|
|
||||||
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):
|
async def no_voice_close_connect(conn, have_voice):
|
||||||
if have_voice:
|
if have_voice:
|
||||||
conn.client_no_voice_last_time = 0.0
|
conn.last_activity_time = time.time() * 1000
|
||||||
return
|
return
|
||||||
if conn.client_no_voice_last_time == 0.0:
|
# 只有在已经初始化过时间戳的情况下才进行超时检查
|
||||||
conn.client_no_voice_last_time = time.time() * 1000
|
if conn.last_activity_time > 0.0:
|
||||||
else:
|
no_voice_time = time.time() * 1000 - conn.last_activity_time
|
||||||
no_voice_time = time.time() * 1000 - conn.client_no_voice_last_time
|
|
||||||
close_connection_no_voice_time = int(
|
close_connection_no_voice_time = int(
|
||||||
conn.config.get("close_connection_no_voice_time", 120)
|
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:
|
if conn.client_abort:
|
||||||
break
|
break
|
||||||
|
|
||||||
# 每分钟重置一次计时器
|
# 重置没有声音的状态
|
||||||
if time.perf_counter() - last_reset_time > 60:
|
conn.last_activity_time = time.time() * 1000
|
||||||
await conn.reset_timeout()
|
|
||||||
last_reset_time = time.perf_counter()
|
|
||||||
|
|
||||||
# 计算预期发送时间
|
# 计算预期发送时间
|
||||||
expected_time = start_time + (play_position / 1000)
|
expected_time = start_time + (play_position / 1000)
|
||||||
|
|||||||
@@ -1,11 +1,11 @@
|
|||||||
import json
|
import json
|
||||||
from core.handle.abortHandle import handleAbortMessage
|
from core.handle.abortHandle import handleAbortMessage
|
||||||
from core.handle.helloHandle import handleHelloMessage
|
from core.handle.helloHandle import handleHelloMessage
|
||||||
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.utils.util import remove_punctuation_and_length, filter_sensitive_info
|
||||||
from core.handle.receiveAudioHandle import startToChat, handleAudioMessage
|
from core.handle.receiveAudioHandle import startToChat, handleAudioMessage
|
||||||
from core.handle.sendAudioHandle import send_stt_message, send_tts_message
|
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
|
from core.handle.reportHandle import enqueue_asr_report
|
||||||
import asyncio
|
import asyncio
|
||||||
|
|
||||||
@@ -77,7 +77,7 @@ async def handleTextMessage(conn, message):
|
|||||||
if "states" in msg_json:
|
if "states" in msg_json:
|
||||||
asyncio.create_task(handleIotStatus(conn, msg_json["states"]))
|
asyncio.create_task(handleIotStatus(conn, msg_json["states"]))
|
||||||
elif msg_json["type"] == "mcp":
|
elif msg_json["type"] == "mcp":
|
||||||
conn.logger.bind(tag=TAG).info(f"收到mcp消息:{message}")
|
conn.logger.bind(tag=TAG).info(f"收到mcp消息:{message[:100]}")
|
||||||
if "payload" in msg_json:
|
if "payload" in msg_json:
|
||||||
asyncio.create_task(
|
asyncio.create_task(
|
||||||
handle_mcp_message(conn, conn.mcp_client, msg_json["payload"])
|
handle_mcp_message(conn, conn.mcp_client, msg_json["payload"])
|
||||||
|
|||||||
@@ -1,130 +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}"
|
|
||||||
)
|
|
||||||
|
|
||||||
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:
|
if "code" in result and result["code"] != 1000:
|
||||||
error_msg = f"ASR服务初始化失败: {result.get('payload_msg', {}).get('message', '未知错误')}"
|
error_msg = f"ASR服务初始化失败: {result.get('payload_msg', {}).get('error', '未知错误')}"
|
||||||
if "payload_msg" in result:
|
|
||||||
error_msg += f"\n详细错误信息: {json.dumps(result['payload_msg'], ensure_ascii=False)}"
|
|
||||||
logger.bind(tag=TAG).error(error_msg)
|
logger.bind(tag=TAG).error(error_msg)
|
||||||
raise Exception(error_msg)
|
raise Exception(error_msg)
|
||||||
|
|
||||||
@@ -256,7 +254,6 @@ class ASRProvider(ASRProviderBase):
|
|||||||
"X-Api-Access-Key": self.access_token,
|
"X-Api-Access-Key": self.access_token,
|
||||||
"X-Api-Resource-Id": "volc.bigasr.sauc.duration",
|
"X-Api-Resource-Id": "volc.bigasr.sauc.duration",
|
||||||
"X-Api-Connect-Id": str(uuid.uuid4()),
|
"X-Api-Connect-Id": str(uuid.uuid4()),
|
||||||
"Host": "openspeech.bytedance.com",
|
|
||||||
}
|
}
|
||||||
|
|
||||||
def generate_header(
|
def generate_header(
|
||||||
@@ -309,9 +306,14 @@ class ASRProvider(ASRProviderBase):
|
|||||||
|
|
||||||
# 如果是错误响应
|
# 如果是错误响应
|
||||||
if message_type == 0x0F: # SERVER_ERROR_RESPONSE
|
if message_type == 0x0F: # SERVER_ERROR_RESPONSE
|
||||||
code = int.from_bytes(header[4:8], "big", signed=False)
|
code = int.from_bytes(res[4:8], "big", signed=False)
|
||||||
error_msg = res[8:].decode("utf-8")
|
msg_length = int.from_bytes(res[8:12], "big", signed=False)
|
||||||
return {"code": code, "error": error_msg}
|
error_msg = json.loads(res[12:].decode("utf-8"))
|
||||||
|
return {
|
||||||
|
"code": code,
|
||||||
|
"msg_length": msg_length,
|
||||||
|
"payload_msg": error_msg,
|
||||||
|
}
|
||||||
|
|
||||||
# 获取JSON数据(跳过12字节头部)
|
# 获取JSON数据(跳过12字节头部)
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -95,6 +95,10 @@ class IntentProvider(IntentProviderBase):
|
|||||||
"1. 只返回JSON格式,不要包含任何其他文字\n"
|
"1. 只返回JSON格式,不要包含任何其他文字\n"
|
||||||
'2. 如果没有找到匹配的函数,返回{"function_call": {"name": "continue_chat"}}\n'
|
'2. 如果没有找到匹配的函数,返回{"function_call": {"name": "continue_chat"}}\n'
|
||||||
"3. 确保返回的JSON格式正确,包含所有必要的字段\n"
|
"3. 确保返回的JSON格式正确,包含所有必要的字段\n"
|
||||||
|
"特殊说明:\n"
|
||||||
|
"- 当用户单次输入包含多个指令时(如'打开灯并且调高音量')\n"
|
||||||
|
"- 请返回多个function_call组成的JSON数组\n"
|
||||||
|
"- 示例:{'function_calls': [{name:'light_on'}, {name:'volume_up'}]}"
|
||||||
)
|
)
|
||||||
return prompt
|
return prompt
|
||||||
|
|
||||||
@@ -172,7 +176,11 @@ class IntentProvider(IntentProviderBase):
|
|||||||
music_file_names = music_config["music_file_names"]
|
music_file_names = music_config["music_file_names"]
|
||||||
prompt_music = f"{self.promot}\n<musicNames>{music_file_names}\n</musicNames>"
|
prompt_music = f"{self.promot}\n<musicNames>{music_file_names}\n</musicNames>"
|
||||||
|
|
||||||
devices = conn.config["plugins"]["home_assistant"].get("devices", [])
|
home_assistant_cfg = conn.config["plugins"].get("home_assistant")
|
||||||
|
if home_assistant_cfg:
|
||||||
|
devices = home_assistant_cfg.get("devices", [])
|
||||||
|
else:
|
||||||
|
devices = []
|
||||||
if len(devices) > 0:
|
if len(devices) > 0:
|
||||||
hass_prompt = "\n下面是我家智能设备列表(位置,设备名,entity_id),可以通过homeassistant控制\n"
|
hass_prompt = "\n下面是我家智能设备列表(位置,设备名,entity_id),可以通过homeassistant控制\n"
|
||||||
for device in devices:
|
for device in devices:
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ import os
|
|||||||
import yaml
|
import yaml
|
||||||
from config.config_loader import get_project_dir
|
from config.config_loader import get_project_dir
|
||||||
from config.manage_api_client import save_mem_local_short
|
from config.manage_api_client import save_mem_local_short
|
||||||
|
from core.utils.util import check_model_key
|
||||||
|
|
||||||
|
|
||||||
short_term_memory_prompt = """
|
short_term_memory_prompt = """
|
||||||
@@ -145,6 +146,10 @@ class MemoryProvider(MemoryProviderBase):
|
|||||||
# 打印使用的模型信息
|
# 打印使用的模型信息
|
||||||
model_info = getattr(self.llm, "model_name", str(self.llm.__class__.__name__))
|
model_info = getattr(self.llm, "model_name", str(self.llm.__class__.__name__))
|
||||||
logger.bind(tag=TAG).debug(f"使用记忆保存模型: {model_info}")
|
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:
|
if self.llm is None:
|
||||||
logger.bind(tag=TAG).error("LLM is not set for memory provider")
|
logger.bind(tag=TAG).error("LLM is not set for memory provider")
|
||||||
return None
|
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)
|
||||||
+33
-32
@@ -1,14 +1,19 @@
|
|||||||
|
"""设备端MCP客户端支持模块"""
|
||||||
|
|
||||||
import json
|
import json
|
||||||
import asyncio
|
import asyncio
|
||||||
|
import re
|
||||||
from concurrent.futures import Future
|
from concurrent.futures import Future
|
||||||
from core.utils.util import get_vision_url, sanitize_tool_name
|
from core.utils.util import get_vision_url, sanitize_tool_name
|
||||||
from core.utils.auth import AuthToken
|
from core.utils.auth import AuthToken
|
||||||
|
from config.logger import setup_logging
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
|
logger = setup_logging()
|
||||||
|
|
||||||
|
|
||||||
class MCPClient:
|
class MCPClient:
|
||||||
"""MCPClient,用于管理MCP状态和工具"""
|
"""设备端MCP客户端,用于管理MCP状态和工具"""
|
||||||
|
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
self.tools = {} # sanitized_name -> tool_data
|
self.tools = {} # sanitized_name -> tool_data
|
||||||
@@ -94,24 +99,24 @@ class MCPClient:
|
|||||||
async def send_mcp_message(conn, payload: dict):
|
async def send_mcp_message(conn, payload: dict):
|
||||||
"""Helper to send MCP messages, encapsulating common logic."""
|
"""Helper to send MCP messages, encapsulating common logic."""
|
||||||
if not conn.features.get("mcp"):
|
if not conn.features.get("mcp"):
|
||||||
conn.logger.bind(tag=TAG).warning("客户端不支持MCP,无法发送MCP消息")
|
logger.bind(tag=TAG).warning("客户端不支持MCP,无法发送MCP消息")
|
||||||
return
|
return
|
||||||
|
|
||||||
message = json.dumps({"type": "mcp", "payload": payload})
|
message = json.dumps({"type": "mcp", "payload": payload})
|
||||||
|
|
||||||
try:
|
try:
|
||||||
await conn.websocket.send(message)
|
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:
|
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):
|
async def handle_mcp_message(conn, mcp_client: MCPClient, payload: dict):
|
||||||
"""处理MCP消息,包括初始化、工具列表和工具调用响应等"""
|
"""处理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):
|
if not isinstance(payload, dict):
|
||||||
conn.logger.bind(tag=TAG).error("MCP消息缺少payload字段或格式错误")
|
logger.bind(tag=TAG).error("MCP消息缺少payload字段或格式错误")
|
||||||
return
|
return
|
||||||
|
|
||||||
# Handle result
|
# Handle result
|
||||||
@@ -121,32 +126,32 @@ async def handle_mcp_message(conn, mcp_client: MCPClient, payload: dict):
|
|||||||
|
|
||||||
# Check for tool call response first
|
# Check for tool call response first
|
||||||
if msg_id in mcp_client.call_results:
|
if msg_id in mcp_client.call_results:
|
||||||
conn.logger.bind(tag=TAG).debug(
|
logger.bind(tag=TAG).debug(
|
||||||
f"收到工具调用响应,ID: {msg_id}, 结果: {result}"
|
f"收到工具调用响应,ID: {msg_id}, 结果: {result}"
|
||||||
)
|
)
|
||||||
await mcp_client.resolve_call_result(msg_id, result)
|
await mcp_client.resolve_call_result(msg_id, result)
|
||||||
return
|
return
|
||||||
|
|
||||||
if msg_id == 1: # mcpInitializeID
|
if msg_id == 1: # mcpInitializeID
|
||||||
conn.logger.bind(tag=TAG).debug("收到MCP初始化响应")
|
logger.bind(tag=TAG).debug("收到MCP初始化响应")
|
||||||
server_info = result.get("serverInfo")
|
server_info = result.get("serverInfo")
|
||||||
if isinstance(server_info, dict):
|
if isinstance(server_info, dict):
|
||||||
name = server_info.get("name")
|
name = server_info.get("name")
|
||||||
version = server_info.get("version")
|
version = server_info.get("version")
|
||||||
conn.logger.bind(tag=TAG).info(
|
logger.bind(tag=TAG).info(
|
||||||
f"客户端MCP服务器信息: name={name}, version={version}"
|
f"客户端MCP服务器信息: name={name}, version={version}"
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
|
|
||||||
elif msg_id == 2: # mcpToolsListID
|
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:
|
if isinstance(result, dict) and "tools" in result:
|
||||||
tools_data = result["tools"]
|
tools_data = result["tools"]
|
||||||
if not isinstance(tools_data, list):
|
if not isinstance(tools_data, list):
|
||||||
conn.logger.bind(tag=TAG).error("工具列表格式错误")
|
logger.bind(tag=TAG).error("工具列表格式错误")
|
||||||
return
|
return
|
||||||
|
|
||||||
conn.logger.bind(tag=TAG).info(
|
logger.bind(tag=TAG).info(
|
||||||
f"客户端设备支持的工具数量: {len(tools_data)}"
|
f"客户端设备支持的工具数量: {len(tools_data)}"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -172,7 +177,7 @@ async def handle_mcp_message(conn, mcp_client: MCPClient, payload: dict):
|
|||||||
"inputSchema": input_schema,
|
"inputSchema": input_schema,
|
||||||
}
|
}
|
||||||
await mcp_client.add_tool(new_tool)
|
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():
|
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", "")
|
next_cursor = result.get("nextCursor", "")
|
||||||
if next_cursor:
|
if next_cursor:
|
||||||
conn.logger.bind(tag=TAG).info(
|
logger.bind(tag=TAG).info(f"有更多工具,nextCursor: {next_cursor}")
|
||||||
f"有更多工具,nextCursor: {next_cursor}"
|
|
||||||
)
|
|
||||||
await send_mcp_tools_list_continue_request(conn, next_cursor)
|
await send_mcp_tools_list_continue_request(conn, next_cursor)
|
||||||
else:
|
else:
|
||||||
await mcp_client.set_ready(True)
|
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
|
return
|
||||||
|
|
||||||
# Handle method calls (requests from the client)
|
# Handle method calls (requests from the client)
|
||||||
elif "method" in payload:
|
elif "method" in payload:
|
||||||
method = payload["method"]
|
method = payload["method"]
|
||||||
conn.logger.bind(tag=TAG).info(f"收到MCP客户端请求: {method}")
|
logger.bind(tag=TAG).info(f"收到MCP客户端请求: {method}")
|
||||||
|
|
||||||
elif "error" in payload:
|
elif "error" in payload:
|
||||||
error_data = payload["error"]
|
error_data = payload["error"]
|
||||||
error_msg = error_data.get("message", "未知错误")
|
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))
|
msg_id = int(payload.get("id", 0))
|
||||||
if msg_id in mcp_client.call_results:
|
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):
|
async def send_mcp_initialize_message(conn):
|
||||||
"""发送MCP初始化消息"""
|
"""发送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)
|
await send_mcp_message(conn, payload)
|
||||||
|
|
||||||
|
|
||||||
@@ -261,7 +266,7 @@ async def send_mcp_tools_list_request(conn):
|
|||||||
"id": 2, # mcpToolsListID
|
"id": 2, # mcpToolsListID
|
||||||
"method": "tools/list",
|
"method": "tools/list",
|
||||||
}
|
}
|
||||||
conn.logger.bind(tag=TAG).debug("发送MCP工具列表请求")
|
logger.bind(tag=TAG).debug("发送MCP工具列表请求")
|
||||||
await send_mcp_message(conn, payload)
|
await send_mcp_message(conn, payload)
|
||||||
|
|
||||||
|
|
||||||
@@ -273,7 +278,7 @@ async def send_mcp_tools_list_continue_request(conn, cursor: str):
|
|||||||
"method": "tools/list",
|
"method": "tools/list",
|
||||||
"params": {"cursor": cursor},
|
"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)
|
await send_mcp_message(conn, payload)
|
||||||
|
|
||||||
|
|
||||||
@@ -307,8 +312,6 @@ async def call_mcp_tool(
|
|||||||
# 如果解析失败,尝试合并多个JSON对象
|
# 如果解析失败,尝试合并多个JSON对象
|
||||||
try:
|
try:
|
||||||
# 使用正则表达式匹配所有JSON对象
|
# 使用正则表达式匹配所有JSON对象
|
||||||
import re
|
|
||||||
|
|
||||||
json_objects = re.findall(r"\{[^{}]*\}", args)
|
json_objects = re.findall(r"\{[^{}]*\}", args)
|
||||||
if len(json_objects) > 1:
|
if len(json_objects) > 1:
|
||||||
# 合并所有JSON对象
|
# 合并所有JSON对象
|
||||||
@@ -327,7 +330,7 @@ async def call_mcp_tool(
|
|||||||
else:
|
else:
|
||||||
raise ValueError(f"参数JSON解析失败: {args}")
|
raise ValueError(f"参数JSON解析失败: {args}")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
conn.logger.bind(tag=TAG).error(
|
logger.bind(tag=TAG).error(
|
||||||
f"参数JSON解析失败: {str(e)}, 原始参数: {args}"
|
f"参数JSON解析失败: {str(e)}, 原始参数: {args}"
|
||||||
)
|
)
|
||||||
raise ValueError(f"参数JSON解析失败: {str(e)}")
|
raise ValueError(f"参数JSON解析失败: {str(e)}")
|
||||||
@@ -353,15 +356,13 @@ async def call_mcp_tool(
|
|||||||
"params": {"name": actual_name, "arguments": arguments},
|
"params": {"name": actual_name, "arguments": arguments},
|
||||||
}
|
}
|
||||||
|
|
||||||
conn.logger.bind(tag=TAG).info(
|
logger.bind(tag=TAG).info(f"发送客户端mcp工具调用请求: {actual_name},参数: {args}")
|
||||||
f"发送客户端mcp工具调用请求: {actual_name},参数: {args}"
|
|
||||||
)
|
|
||||||
await send_mcp_message(conn, payload)
|
await send_mcp_message(conn, payload)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# Wait for response or timeout
|
# Wait for response or timeout
|
||||||
raw_result = await asyncio.wait_for(result_future, timeout=timeout)
|
raw_result = await asyncio.wait_for(result_future, timeout=timeout)
|
||||||
conn.logger.bind(tag=TAG).info(
|
logger.bind(tag=TAG).info(
|
||||||
f"客户端mcp工具调用 {actual_name} 成功,原始结果: {raw_result}"
|
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"]
|
||||||
+76
-11
@@ -1,7 +1,12 @@
|
|||||||
|
"""服务端MCP客户端"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from datetime import timedelta
|
from datetime import timedelta
|
||||||
import asyncio, os, shutil, concurrent.futures
|
import asyncio
|
||||||
|
import os
|
||||||
|
import shutil
|
||||||
|
import concurrent.futures
|
||||||
from contextlib import AsyncExitStack
|
from contextlib import AsyncExitStack
|
||||||
from typing import Optional, List, Dict, Any
|
from typing import Optional, List, Dict, Any
|
||||||
|
|
||||||
@@ -14,8 +19,15 @@ from core.utils.util import sanitize_tool_name
|
|||||||
TAG = __name__
|
TAG = __name__
|
||||||
|
|
||||||
|
|
||||||
class MCPClient:
|
class ServerMCPClient:
|
||||||
|
"""服务端MCP客户端,用于连接和管理MCP服务"""
|
||||||
|
|
||||||
def __init__(self, config: Dict[str, Any]):
|
def __init__(self, config: Dict[str, Any]):
|
||||||
|
"""初始化服务端MCP客户端
|
||||||
|
|
||||||
|
Args:
|
||||||
|
config: MCP服务配置字典
|
||||||
|
"""
|
||||||
self.logger = setup_logging()
|
self.logger = setup_logging()
|
||||||
self.config = config
|
self.config = config
|
||||||
|
|
||||||
@@ -24,21 +36,26 @@ class MCPClient:
|
|||||||
self._shutdown_evt = asyncio.Event()
|
self._shutdown_evt = asyncio.Event()
|
||||||
|
|
||||||
self.session: Optional[ClientSession] = None
|
self.session: Optional[ClientSession] = None
|
||||||
self.tools: List = [] # original tool objects
|
self.tools: List = [] # 原始工具对象
|
||||||
self.tools_dict: Dict[str, Any] = {}
|
self.tools_dict: Dict[str, Any] = {}
|
||||||
self.name_mapping: Dict[str, str] = {}
|
self.name_mapping: Dict[str, str] = {}
|
||||||
|
|
||||||
async def initialize(self):
|
async def initialize(self):
|
||||||
|
"""初始化MCP客户端连接"""
|
||||||
if self._worker_task:
|
if self._worker_task:
|
||||||
return
|
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()
|
await self._ready_evt.wait()
|
||||||
|
|
||||||
self.logger.bind(tag=TAG).info(
|
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):
|
async def cleanup(self):
|
||||||
|
"""清理MCP客户端资源"""
|
||||||
if not self._worker_task:
|
if not self._worker_task:
|
||||||
return
|
return
|
||||||
|
|
||||||
@@ -46,14 +63,27 @@ class MCPClient:
|
|||||||
try:
|
try:
|
||||||
await asyncio.wait_for(self._worker_task, timeout=20)
|
await asyncio.wait_for(self._worker_task, timeout=20)
|
||||||
except (asyncio.TimeoutError, Exception) as e:
|
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:
|
finally:
|
||||||
self._worker_task = None
|
self._worker_task = None
|
||||||
|
|
||||||
def has_tool(self, name: str) -> bool:
|
def has_tool(self, name: str) -> bool:
|
||||||
|
"""检查是否包含指定工具
|
||||||
|
|
||||||
|
Args:
|
||||||
|
name: 工具名称
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
bool: 是否包含该工具
|
||||||
|
"""
|
||||||
return name in self.tools_dict
|
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 [
|
return [
|
||||||
{
|
{
|
||||||
"type": "function",
|
"type": "function",
|
||||||
@@ -66,9 +96,21 @@ class MCPClient:
|
|||||||
for name, tool in self.tools_dict.items()
|
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:
|
if not self.session:
|
||||||
raise RuntimeError("MCPClient not initialized")
|
raise RuntimeError("服务端MCP客户端未初始化")
|
||||||
|
|
||||||
real_name = self.name_mapping.get(name, name)
|
real_name = self.name_mapping.get(name, name)
|
||||||
loop = self._worker_task.get_loop()
|
loop = self._worker_task.get_loop()
|
||||||
@@ -80,7 +122,29 @@ class MCPClient:
|
|||||||
fut: concurrent.futures.Future = asyncio.run_coroutine_threadsafe(coro, loop)
|
fut: concurrent.futures.Future = asyncio.run_coroutine_threadsafe(coro, loop)
|
||||||
return await asyncio.wrap_future(fut)
|
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):
|
async def _worker(self):
|
||||||
|
"""MCP客户端工作协程"""
|
||||||
async with AsyncExitStack() as stack:
|
async with AsyncExitStack() as stack:
|
||||||
try:
|
try:
|
||||||
# 建立 StdioClient
|
# 建立 StdioClient
|
||||||
@@ -100,6 +164,7 @@ class MCPClient:
|
|||||||
stdio_client(params)
|
stdio_client(params)
|
||||||
)
|
)
|
||||||
read_stream, write_stream = stdio_r, stdio_w
|
read_stream, write_stream = stdio_r, stdio_w
|
||||||
|
|
||||||
# 建立SSEClient
|
# 建立SSEClient
|
||||||
elif "url" in self.config:
|
elif "url" in self.config:
|
||||||
if "API_ACCESS_TOKEN" in self.config:
|
if "API_ACCESS_TOKEN" in self.config:
|
||||||
@@ -114,7 +179,7 @@ class MCPClient:
|
|||||||
read_stream, write_stream = sse_r, sse_w
|
read_stream, write_stream = sse_r, sse_w
|
||||||
|
|
||||||
else:
|
else:
|
||||||
raise ValueError("MCPClient config must include 'command' or 'url'")
|
raise ValueError("MCP客户端配置必须包含'command'或'url'")
|
||||||
|
|
||||||
self.session = await stack.enter_async_context(
|
self.session = await stack.enter_async_context(
|
||||||
ClientSession(
|
ClientSession(
|
||||||
@@ -138,6 +203,6 @@ class MCPClient:
|
|||||||
await self._shutdown_evt.wait()
|
await self._shutdown_evt.wait()
|
||||||
|
|
||||||
except Exception as e:
|
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()
|
self._ready_evt.set()
|
||||||
raise
|
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.tts_text_buff = []
|
||||||
self.punctuations = (
|
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:
|
for seg in segments:
|
||||||
self.tts_text_queue.put(
|
self.tts_text_queue.put(
|
||||||
TTSMessageDTO(
|
TTSMessageDTO(
|
||||||
|
|||||||
@@ -12,6 +12,8 @@ from core.utils.util import check_model_key
|
|||||||
from core.providers.tts.base import TTSProviderBase
|
from core.providers.tts.base import TTSProviderBase
|
||||||
from core.handle.abortHandle import handleAbortMessage
|
from core.handle.abortHandle import handleAbortMessage
|
||||||
from core.providers.tts.dto.dto import SentenceType, ContentType, InterfaceType
|
from core.providers.tts.dto.dto import SentenceType, ContentType, InterfaceType
|
||||||
|
from asyncio import Task
|
||||||
|
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
logger = setup_logging()
|
||||||
@@ -141,6 +143,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
super().__init__(config, delete_audio_file)
|
super().__init__(config, delete_audio_file)
|
||||||
self.ws = None
|
self.ws = None
|
||||||
self.interface_type = InterfaceType.DUAL_STREAM
|
self.interface_type = InterfaceType.DUAL_STREAM
|
||||||
|
self._monitor_task = None # 监听任务引用
|
||||||
self.appId = config.get("appid")
|
self.appId = config.get("appid")
|
||||||
self.access_token = config.get("access_token")
|
self.access_token = config.get("access_token")
|
||||||
self.cluster = config.get("cluster")
|
self.cluster = config.get("cluster")
|
||||||
@@ -270,8 +273,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
try:
|
try:
|
||||||
# 建立新连接
|
# 建立新连接
|
||||||
if self.ws is None:
|
if self.ws is None:
|
||||||
await handleAbortMessage(self.conn)
|
logger.bind(tag=TAG).warning(f"WebSocket连接不存在,终止发送文本")
|
||||||
logger.bind(tag=TAG).error(f"WebSocket连接不存在,终止发送文本")
|
|
||||||
return
|
return
|
||||||
|
|
||||||
# 过滤Markdown
|
# 过滤Markdown
|
||||||
@@ -293,6 +295,25 @@ class TTSProvider(TTSProviderBase):
|
|||||||
async def start_session(self, session_id):
|
async def start_session(self, session_id):
|
||||||
logger.bind(tag=TAG).info(f"开始会话~~{session_id}")
|
logger.bind(tag=TAG).info(f"开始会话~~{session_id}")
|
||||||
try:
|
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()
|
await self._ensure_connection()
|
||||||
|
|
||||||
@@ -463,6 +484,8 @@ class TTSProvider(TTSProviderBase):
|
|||||||
except:
|
except:
|
||||||
pass
|
pass
|
||||||
self.ws = None
|
self.ws = None
|
||||||
|
# 监听任务退出时清理引用
|
||||||
|
self._monitor_task = None
|
||||||
|
|
||||||
async def send_event(
|
async def send_event(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -35,6 +35,10 @@ class VADProvider(VADProviderBase):
|
|||||||
pcm_frame = self.decoder.decode(opus_packet, 960)
|
pcm_frame = self.decoder.decode(opus_packet, 960)
|
||||||
conn.client_audio_buffer.extend(pcm_frame) # 将新数据加入缓冲区
|
conn.client_audio_buffer.extend(pcm_frame) # 将新数据加入缓冲区
|
||||||
|
|
||||||
|
# 确保帧计数器存在
|
||||||
|
if not hasattr(conn, "client_voice_frame_count"):
|
||||||
|
conn.client_voice_frame_count = 0
|
||||||
|
|
||||||
# 处理缓冲区中的完整帧(每次处理512采样点)
|
# 处理缓冲区中的完整帧(每次处理512采样点)
|
||||||
client_have_voice = False
|
client_have_voice = False
|
||||||
while len(conn.client_audio_buffer) >= 512 * 2:
|
while len(conn.client_audio_buffer) >= 512 * 2:
|
||||||
@@ -50,18 +54,24 @@ class VADProvider(VADProviderBase):
|
|||||||
# 检测语音活动
|
# 检测语音活动
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
speech_prob = self.model(audio_tensor, 16000).item()
|
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:
|
if conn.client_have_voice and not client_have_voice:
|
||||||
stop_duration = (
|
stop_duration = time.time() * 1000 - conn.last_activity_time
|
||||||
time.time() * 1000 - conn.client_have_voice_last_time
|
|
||||||
)
|
|
||||||
if stop_duration >= self.silence_threshold_ms:
|
if stop_duration >= self.silence_threshold_ms:
|
||||||
conn.client_voice_stop = True
|
conn.client_voice_stop = True
|
||||||
if client_have_voice:
|
if client_have_voice:
|
||||||
conn.client_have_voice = True
|
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
|
return client_have_voice
|
||||||
except opuslib_next.OpusError as e:
|
except opuslib_next.OpusError as e:
|
||||||
|
|||||||
@@ -52,7 +52,7 @@ services:
|
|||||||
volumes:
|
volumes:
|
||||||
# 配置文件目录
|
# 配置文件目录
|
||||||
- ./uploadfile:/uploadfile
|
- ./uploadfile:/uploadfile
|
||||||
|
# 数据库模块
|
||||||
xiaozhi-esp32-server-db:
|
xiaozhi-esp32-server-db:
|
||||||
image: mysql:latest
|
image: mysql:latest
|
||||||
container_name: xiaozhi-esp32-server-db
|
container_name: xiaozhi-esp32-server-db
|
||||||
@@ -73,6 +73,7 @@ services:
|
|||||||
- MYSQL_ROOT_PASSWORD=123456
|
- MYSQL_ROOT_PASSWORD=123456
|
||||||
- MYSQL_DATABASE=xiaozhi_esp32_server
|
- MYSQL_DATABASE=xiaozhi_esp32_server
|
||||||
- MYSQL_INITDB_ARGS="--character-set-server=utf8mb4 --collation-server=utf8mb4_unicode_ci"
|
- MYSQL_INITDB_ARGS="--character-set-server=utf8mb4 --collation-server=utf8mb4_unicode_ci"
|
||||||
|
# redis模块
|
||||||
xiaozhi-esp32-server-redis:
|
xiaozhi-esp32-server-redis:
|
||||||
image: redis
|
image: redis
|
||||||
expose:
|
expose:
|
||||||
|
|||||||
@@ -31,7 +31,8 @@ hass_set_state_function_desc = {
|
|||||||
"description": "只有在设置静音操作时才需要,设置静音的时候该值为true,取消静音时该值为false",
|
"description": "只有在设置静音操作时才需要,设置静音的时候该值为true,取消静音时该值为false",
|
||||||
},
|
},
|
||||||
"rgb_color": {
|
"rgb_color": {
|
||||||
"type": "list",
|
"type": "array",
|
||||||
|
"items": {"type": "integer"},
|
||||||
"description": "只有在设置颜色时需要,这里填目标颜色的rgb值",
|
"description": "只有在设置颜色时需要,这里填目标颜色的rgb值",
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
|
|||||||
@@ -35,7 +35,7 @@ class Action(Enum):
|
|||||||
|
|
||||||
|
|
||||||
class ActionResponse:
|
class ActionResponse:
|
||||||
def __init__(self, action: Action, result, response):
|
def __init__(self, action: Action, result=None, response=None):
|
||||||
self.action = action # 动作类型
|
self.action = action # 动作类型
|
||||||
self.result = result # 动作产生的结果
|
self.result = result # 动作产生的结果
|
||||||
self.response = response # 直接回复的内容
|
self.response = response # 直接回复的内容
|
||||||
|
|||||||
@@ -30,7 +30,7 @@ baidu-aip==4.16.13
|
|||||||
chardet==5.2.0
|
chardet==5.2.0
|
||||||
aioconsole==0.8.1
|
aioconsole==0.8.1
|
||||||
markitdown==0.1.1
|
markitdown==0.1.1
|
||||||
mcp-proxy==0.6.0
|
mcp-proxy==0.8.0
|
||||||
PyJWT==2.8.0
|
PyJWT==2.8.0
|
||||||
psutil==7.0.0
|
psutil==7.0.0
|
||||||
portalocker==2.10.1
|
portalocker==2.10.1
|
||||||
Reference in New Issue
Block a user