mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-27 01:23:55 +08:00
update:完成mcp接入点对接
This commit is contained in:
@@ -104,7 +104,8 @@ wakeup_words:
|
|||||||
- "喵喵同学"
|
- "喵喵同学"
|
||||||
- "小滨小滨"
|
- "小滨小滨"
|
||||||
- "小冰小冰"
|
- "小冰小冰"
|
||||||
|
# MCP接入点地址
|
||||||
|
mcp_endpoint: 你的接入点 websocket地址
|
||||||
# 插件的基础配置
|
# 插件的基础配置
|
||||||
plugins:
|
plugins:
|
||||||
# 获取天气插件的配置,这里填写你的api_key
|
# 获取天气插件的配置,这里填写你的api_key
|
||||||
|
|||||||
@@ -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"])
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ class ToolType(Enum):
|
|||||||
SERVER_MCP = "server_mcp" # 服务端MCP
|
SERVER_MCP = "server_mcp" # 服务端MCP
|
||||||
DEVICE_IOT = "device_iot" # 设备端IoT
|
DEVICE_IOT = "device_iot" # 设备端IoT
|
||||||
DEVICE_MCP = "device_mcp" # 设备端MCP
|
DEVICE_MCP = "device_mcp" # 设备端MCP
|
||||||
|
MCP_ENDPOINT = "mcp_endpoint" # MCP接入点
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
|
|||||||
@@ -113,7 +113,7 @@ async def send_mcp_message(conn, payload: dict):
|
|||||||
|
|
||||||
async def handle_mcp_message(conn, mcp_client: MCPClient, payload: dict):
|
async def handle_mcp_message(conn, mcp_client: MCPClient, payload: dict):
|
||||||
"""处理MCP消息,包括初始化、工具列表和工具调用响应等"""
|
"""处理MCP消息,包括初始化、工具列表和工具调用响应等"""
|
||||||
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):
|
||||||
logger.bind(tag=TAG).error("MCP消息缺少payload字段或格式错误")
|
logger.bind(tag=TAG).error("MCP消息缺少payload字段或格式错误")
|
||||||
|
|||||||
@@ -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,111 @@
|
|||||||
|
"""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):
|
||||||
|
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,365 @@
|
|||||||
|
"""MCP接入点处理器"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
import asyncio
|
||||||
|
import re
|
||||||
|
import websockets
|
||||||
|
from concurrent.futures import Future
|
||||||
|
from core.utils.util import sanitize_tool_name
|
||||||
|
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) -> MCPEndpointClient:
|
||||||
|
"""连接到MCP接入点"""
|
||||||
|
if not mcp_endpoint_url:
|
||||||
|
logger.bind(tag=TAG).warning("MCP接入点URL为空,跳过连接")
|
||||||
|
return None
|
||||||
|
|
||||||
|
try:
|
||||||
|
logger.bind(tag=TAG).info(f"正在连接到MCP接入点: {mcp_endpoint_url}")
|
||||||
|
websocket = await websockets.connect(mcp_endpoint_url)
|
||||||
|
|
||||||
|
mcp_client = MCPEndpointClient()
|
||||||
|
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接入点初始化响应")
|
||||||
|
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}"
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
elif msg_id == 2: # mcpToolsListID
|
||||||
|
logger.bind(tag=TAG).debug("收到MCP接入点工具列表响应")
|
||||||
|
if 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 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接入点工具已获取,客户端准备就绪"
|
||||||
|
)
|
||||||
|
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
|
||||||
@@ -1,4 +1,4 @@
|
|||||||
"""统一工具处理器,替换原有的FunctionHandler"""
|
"""统一工具处理器"""
|
||||||
|
|
||||||
import json
|
import json
|
||||||
from typing import Dict, List, Any, Optional
|
from typing import Dict, List, Any, Optional
|
||||||
@@ -12,6 +12,7 @@ from .server_plugins import ServerPluginExecutor
|
|||||||
from .server_mcp import ServerMCPExecutor
|
from .server_mcp import ServerMCPExecutor
|
||||||
from .device_iot import DeviceIoTExecutor
|
from .device_iot import DeviceIoTExecutor
|
||||||
from .device_mcp import DeviceMCPExecutor
|
from .device_mcp import DeviceMCPExecutor
|
||||||
|
from .mcp_endpoint import MCPEndpointExecutor
|
||||||
|
|
||||||
|
|
||||||
class UnifiedToolHandler:
|
class UnifiedToolHandler:
|
||||||
@@ -30,6 +31,7 @@ class UnifiedToolHandler:
|
|||||||
self.server_mcp_executor = ServerMCPExecutor(conn)
|
self.server_mcp_executor = ServerMCPExecutor(conn)
|
||||||
self.device_iot_executor = DeviceIoTExecutor(conn)
|
self.device_iot_executor = DeviceIoTExecutor(conn)
|
||||||
self.device_mcp_executor = DeviceMCPExecutor(conn)
|
self.device_mcp_executor = DeviceMCPExecutor(conn)
|
||||||
|
self.mcp_endpoint_executor = MCPEndpointExecutor(conn)
|
||||||
|
|
||||||
# 注册执行器
|
# 注册执行器
|
||||||
self.tool_manager.register_executor(
|
self.tool_manager.register_executor(
|
||||||
@@ -44,6 +46,9 @@ class UnifiedToolHandler:
|
|||||||
self.tool_manager.register_executor(
|
self.tool_manager.register_executor(
|
||||||
ToolType.DEVICE_MCP, self.device_mcp_executor
|
ToolType.DEVICE_MCP, self.device_mcp_executor
|
||||||
)
|
)
|
||||||
|
self.tool_manager.register_executor(
|
||||||
|
ToolType.MCP_ENDPOINT, self.mcp_endpoint_executor
|
||||||
|
)
|
||||||
|
|
||||||
# 初始化标志
|
# 初始化标志
|
||||||
self.finish_init = False
|
self.finish_init = False
|
||||||
@@ -57,6 +62,9 @@ class UnifiedToolHandler:
|
|||||||
# 初始化服务端MCP
|
# 初始化服务端MCP
|
||||||
await self.server_mcp_executor.initialize()
|
await self.server_mcp_executor.initialize()
|
||||||
|
|
||||||
|
# 初始化MCP接入点
|
||||||
|
await self._initialize_mcp_endpoint()
|
||||||
|
|
||||||
# 初始化Home Assistant(如果需要)
|
# 初始化Home Assistant(如果需要)
|
||||||
self._initialize_home_assistant()
|
self._initialize_home_assistant()
|
||||||
|
|
||||||
@@ -66,6 +74,28 @@ class UnifiedToolHandler:
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
self.logger.error(f"统一工具处理器初始化失败: {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:
|
||||||
|
self.logger.info(f"正在初始化MCP接入点: {mcp_endpoint_url}")
|
||||||
|
mcp_endpoint_client = await connect_mcp_endpoint(mcp_endpoint_url)
|
||||||
|
|
||||||
|
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):
|
def _initialize_home_assistant(self):
|
||||||
"""初始化Home Assistant提示词"""
|
"""初始化Home Assistant提示词"""
|
||||||
try:
|
try:
|
||||||
@@ -183,6 +213,14 @@ class UnifiedToolHandler:
|
|||||||
"""清理资源"""
|
"""清理资源"""
|
||||||
try:
|
try:
|
||||||
await self.server_mcp_executor.cleanup()
|
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("工具处理器清理完成")
|
self.logger.info("工具处理器清理完成")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
self.logger.error(f"工具处理器清理失败: {e}")
|
self.logger.error(f"工具处理器清理失败: {e}")
|
||||||
|
|||||||
Reference in New Issue
Block a user