mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-31 03:43:57 +08:00
update:优化设备端读取mcp接入点工具
This commit is contained in:
+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(
|
||||||
|
|||||||
@@ -473,6 +473,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,
|
||||||
|
|||||||
@@ -58,10 +58,7 @@ async def handleIotDescriptors(conn, descriptors):
|
|||||||
# 注册IoT工具到统一工具处理器
|
# 注册IoT工具到统一工具处理器
|
||||||
await conn.func_handler.register_iot_tools(descriptors)
|
await conn.func_handler.register_iot_tools(descriptors)
|
||||||
|
|
||||||
func_names = conn.func_handler.current_support_functions()
|
conn.func_handler.current_support_functions()
|
||||||
logger.bind(tag=TAG).info(
|
|
||||||
f"更新function描述列表完成,当前支持的函数: {func_names}"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
async def handleIotStatus(conn, states):
|
async def handleIotStatus(conn, states):
|
||||||
|
|||||||
@@ -200,6 +200,11 @@ async def handle_mcp_message(conn, mcp_client: MCPClient, payload: dict):
|
|||||||
else:
|
else:
|
||||||
await mcp_client.set_ready(True)
|
await mcp_client.set_ready(True)
|
||||||
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)
|
||||||
|
|||||||
@@ -12,7 +12,8 @@ logger = setup_logging()
|
|||||||
class MCPEndpointClient:
|
class MCPEndpointClient:
|
||||||
"""MCP接入点客户端,用于管理MCP接入点状态和工具"""
|
"""MCP接入点客户端,用于管理MCP接入点状态和工具"""
|
||||||
|
|
||||||
def __init__(self):
|
def __init__(self, conn=None):
|
||||||
|
self.conn = conn
|
||||||
self.tools = {} # sanitized_name -> tool_data
|
self.tools = {} # sanitized_name -> tool_data
|
||||||
self.name_mapping = {}
|
self.name_mapping = {}
|
||||||
self.ready = False
|
self.ready = False
|
||||||
|
|||||||
@@ -4,8 +4,6 @@ import json
|
|||||||
import asyncio
|
import asyncio
|
||||||
import re
|
import re
|
||||||
import websockets
|
import websockets
|
||||||
from concurrent.futures import Future
|
|
||||||
from core.utils.util import sanitize_tool_name
|
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
from .mcp_endpoint_client import MCPEndpointClient
|
from .mcp_endpoint_client import MCPEndpointClient
|
||||||
|
|
||||||
@@ -13,16 +11,15 @@ TAG = __name__
|
|||||||
logger = setup_logging()
|
logger = setup_logging()
|
||||||
|
|
||||||
|
|
||||||
async def connect_mcp_endpoint(mcp_endpoint_url: str) -> MCPEndpointClient:
|
async def connect_mcp_endpoint(mcp_endpoint_url: str, conn=None) -> MCPEndpointClient:
|
||||||
"""连接到MCP接入点"""
|
"""连接到MCP接入点"""
|
||||||
if not mcp_endpoint_url or "你的" in mcp_endpoint_url or mcp_endpoint_url == "null":
|
if not mcp_endpoint_url or "你的" in mcp_endpoint_url or mcp_endpoint_url == "null":
|
||||||
return None
|
return None
|
||||||
|
|
||||||
try:
|
try:
|
||||||
logger.bind(tag=TAG).info(f"正在连接到MCP接入点: {mcp_endpoint_url}")
|
|
||||||
websocket = await websockets.connect(mcp_endpoint_url)
|
websocket = await websockets.connect(mcp_endpoint_url)
|
||||||
|
|
||||||
mcp_client = MCPEndpointClient()
|
mcp_client = MCPEndpointClient(conn)
|
||||||
mcp_client.set_websocket(websocket)
|
mcp_client.set_websocket(websocket)
|
||||||
|
|
||||||
# 启动消息监听器
|
# 启动消息监听器
|
||||||
@@ -85,18 +82,27 @@ async def handle_mcp_endpoint_message(mcp_client: MCPEndpointClient, message: st
|
|||||||
|
|
||||||
if msg_id == 1: # mcpInitializeID
|
if msg_id == 1: # mcpInitializeID
|
||||||
logger.bind(tag=TAG).debug("收到MCP接入点初始化响应")
|
logger.bind(tag=TAG).debug("收到MCP接入点初始化响应")
|
||||||
server_info = result.get("serverInfo")
|
if result is not None and isinstance(result, dict):
|
||||||
if isinstance(server_info, dict):
|
server_info = result.get("serverInfo")
|
||||||
name = server_info.get("name")
|
if isinstance(server_info, dict):
|
||||||
version = server_info.get("version")
|
name = server_info.get("name")
|
||||||
logger.bind(tag=TAG).info(
|
version = server_info.get("version")
|
||||||
f"MCP接入点服务器信息: name={name}, version={version}"
|
logger.bind(tag=TAG).info(
|
||||||
|
f"MCP接入点服务器信息: name={name}, version={version}"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
logger.bind(tag=TAG).warning(
|
||||||
|
"MCP接入点初始化响应结果为空或格式错误"
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
|
|
||||||
elif msg_id == 2: # mcpToolsListID
|
elif msg_id == 2: # mcpToolsListID
|
||||||
logger.bind(tag=TAG).debug("收到MCP接入点工具列表响应")
|
logger.bind(tag=TAG).debug("收到MCP接入点工具列表响应")
|
||||||
if isinstance(result, dict) and "tools" in result:
|
if (
|
||||||
|
result is not None
|
||||||
|
and 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):
|
||||||
logger.bind(tag=TAG).error("工具列表格式错误")
|
logger.bind(tag=TAG).error("工具列表格式错误")
|
||||||
@@ -152,7 +158,9 @@ async def handle_mcp_endpoint_message(mcp_client: MCPEndpointClient, message: st
|
|||||||
)
|
)
|
||||||
tool_data["description"] = description
|
tool_data["description"] = description
|
||||||
|
|
||||||
next_cursor = result.get("nextCursor", "")
|
next_cursor = (
|
||||||
|
result.get("nextCursor", "") if result is not None else ""
|
||||||
|
)
|
||||||
if next_cursor:
|
if next_cursor:
|
||||||
logger.bind(tag=TAG).info(
|
logger.bind(tag=TAG).info(
|
||||||
f"有更多工具,nextCursor: {next_cursor}"
|
f"有更多工具,nextCursor: {next_cursor}"
|
||||||
@@ -165,6 +173,24 @@ async def handle_mcp_endpoint_message(mcp_client: MCPEndpointClient, message: st
|
|||||||
logger.bind(tag=TAG).info(
|
logger.bind(tag=TAG).info(
|
||||||
"所有MCP接入点工具已获取,客户端准备就绪"
|
"所有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
|
return
|
||||||
|
|
||||||
# Handle method calls (requests from the endpoint)
|
# Handle method calls (requests from the endpoint)
|
||||||
|
|||||||
@@ -66,6 +66,10 @@ class ServerMCPManager:
|
|||||||
f"Failed to initialize MCP server {name}: {e}"
|
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]]:
|
def get_all_tools(self) -> List[Dict[str, Any]]:
|
||||||
"""获取所有服务的工具function定义"""
|
"""获取所有服务的工具function定义"""
|
||||||
return self.tools
|
return self.tools
|
||||||
|
|||||||
@@ -71,6 +71,9 @@ class UnifiedToolHandler:
|
|||||||
self.finish_init = True
|
self.finish_init = True
|
||||||
self.logger.info("统一工具处理器初始化完成")
|
self.logger.info("统一工具处理器初始化完成")
|
||||||
|
|
||||||
|
# 输出当前支持的所有工具列表
|
||||||
|
self.current_support_functions()
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
self.logger.error(f"统一工具处理器初始化失败: {e}")
|
self.logger.error(f"统一工具处理器初始化失败: {e}")
|
||||||
|
|
||||||
@@ -88,7 +91,9 @@ class UnifiedToolHandler:
|
|||||||
and mcp_endpoint_url != "null"
|
and mcp_endpoint_url != "null"
|
||||||
):
|
):
|
||||||
self.logger.info(f"正在初始化MCP接入点: {mcp_endpoint_url}")
|
self.logger.info(f"正在初始化MCP接入点: {mcp_endpoint_url}")
|
||||||
mcp_endpoint_client = await connect_mcp_endpoint(mcp_endpoint_url)
|
mcp_endpoint_client = await connect_mcp_endpoint(
|
||||||
|
mcp_endpoint_url, self.conn
|
||||||
|
)
|
||||||
|
|
||||||
if mcp_endpoint_client:
|
if mcp_endpoint_client:
|
||||||
# 将MCP接入点客户端保存到连接对象中
|
# 将MCP接入点客户端保存到连接对象中
|
||||||
|
|||||||
Reference in New Issue
Block a user