import json import asyncio from concurrent.futures import Future TAG = __name__ class MCPClient: """MCPClient,用于管理MCP状态和工具""" def __init__(self): self.tools = [] self.ready = False self.call_results = {} # To store Futures for tool call responses self.next_id = 1 self.lock = asyncio.Lock() async def has_tool(self, name: str) -> bool: async with self.lock: for tool in self.tools: if tool["name"] == name: return True return False def get_available_tools(self) -> list: # async with self.lock: result = [] for tool in self.tools: function_def = { "name": tool["name"], "description": tool["description"], "parameters": { "type": tool["inputSchema"].get("type", "object"), "properties": tool["inputSchema"].get("properties", {}), "required": tool["inputSchema"].get("required", []) } } result.append({"type": "function", "function": function_def}) 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: self.tools.append(tool_data) 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) async def send_mcp_message(conn, payload: dict): """Helper to send MCP messages, encapsulating common logic.""" if not conn.features.get("mcp"): conn.logger.bind(tag=TAG).warning("客户端不支持MCP,无法发送MCP消息") return message = json.dumps({ "type": "mcp", "payload": payload }) try: await conn.websocket.send(message) conn.logger.bind(tag=TAG).info(f"成功发送MCP消息: {message}") except Exception as e: conn.logger.bind(tag=TAG).error(f"发送MCP消息失败: {e}") async def handle_mcp_message(conn, mcp_client: MCPClient, payload: dict): """处理MCP消息,包括初始化、工具列表和工具调用响应等""" conn.logger.bind(tag=TAG).info(f"处理MCP消息: {payload}") if not isinstance(payload, dict): conn.logger.bind(tag=TAG).error("MCP消息缺少payload字段或格式错误") return # Handle result if "result" in payload: result = payload["result"] msg_id = int(payload.get("id", 0)) # Check for tool call response first if msg_id in mcp_client.call_results: conn.logger.bind(tag=TAG).debug(f"收到工具调用响应,ID: {msg_id}, 结果: {result}") await mcp_client.resolve_call_result(msg_id, result) return if msg_id == 1: # mcpInitializeID conn.logger.bind(tag=TAG).debug("收到MCP初始化响应") server_info = result.get("serverInfo") if isinstance(server_info, dict): name = server_info.get("name") version = server_info.get("version") conn.logger.bind(tag=TAG).info(f"客户端MCP服务器信息: name={name}, version={version}") await send_mcp_tools_list_request(conn) # After initialization, request tool list return elif msg_id == 2: # mcpToolsListID conn.logger.bind(tag=TAG).debug("收到MCP工具列表响应") if isinstance(result, dict) and "tools" in result: tools_data = result["tools"] if not isinstance(tools_data, list): conn.logger.bind(tag=TAG).error("工具列表格式错误") return conn.logger.bind(tag=TAG).info(f"客户端设备支持的工具数量: {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) conn.logger.bind(tag=TAG).debug(f"客户端工具 #{i+1}: {name}") next_cursor = result.get("nextCursor", "") if next_cursor: conn.logger.bind(tag=TAG).info(f"有更多工具,nextCursor: {next_cursor}") await send_mcp_tools_list_continue_request(conn, next_cursor) else: await mcp_client.set_ready(True) conn.logger.bind(tag=TAG).info("所有工具已获取,MCP客户端准备就绪") return # Handle method calls (requests from the client) elif "method" in payload: method = payload["method"] conn.logger.bind(tag=TAG).info(f"收到MCP客户端请求: {method}") elif "error" in payload: error_data = payload["error"] error_msg = error_data.get("message", "未知错误") conn.logger.bind(tag=TAG).error(f"收到MCP错误响应: {error_msg}") msg_id = int(payload.get("id", 0)) if msg_id in mcp_client.call_results: await mcp_client.reject_call_result(msg_id, Exception(f"MCP错误: {error_msg}")) # --- Outgoing MCP Messages --- async def send_mcp_initialize_message(conn): """发送MCP初始化消息""" payload = { "jsonrpc": "2.0", "id": 1, # mcpInitializeID "method": "initialize", "params": { "protocolVersion": "2024-11-05", "capabilities": { "roots": {"listChanged": True}, "sampling": {}, }, "clientInfo": { "name": "XiaozhiClient", "version": "1.0.0", }, }, } conn.logger.bind(tag=TAG).info("发送MCP初始化消息") await send_mcp_message(conn, payload) async def send_mcp_tools_list_request(conn): """发送MCP工具列表请求""" payload = { "jsonrpc": "2.0", "id": 2, # mcpToolsListID "method": "tools/list", } conn.logger.bind(tag=TAG).debug("发送MCP工具列表请求") await send_mcp_message(conn, payload) async def send_mcp_tools_list_continue_request(conn, cursor: str): """发送带有cursor的MCP工具列表请求""" payload = { "jsonrpc": "2.0", "id": 2, # mcpToolsListID (same ID for continuation) "method": "tools/list", "params": {"cursor": cursor}, } conn.logger.bind(tag=TAG).info(f"发送带cursor的MCP工具列表请求: {cursor}") await send_mcp_message(conn, payload) async def call_mcp_tool(conn, mcp_client: MCPClient, tool_name: str, args: str = '{}', timeout: int = 30): """ 调用指定的工具,并等待响应 """ if not await mcp_client.is_ready(): raise RuntimeError("MCP客户端尚未准备就绪") if not await 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) payload = { "jsonrpc": "2.0", "id": tool_call_id, "method": "tools/call", "params": { "name": tool_name, "arguments": json.loads(args) if isinstance(args, str) else args }, } conn.logger.bind(tag=TAG).info(f"发送客户端mcp工具调用请求: {tool_name},参数: {args}") await send_mcp_message(conn, payload) try: # Wait for response or timeout raw_result = await asyncio.wait_for(result_future, timeout=timeout) conn.logger.bind(tag=TAG).info(f"客户端mcp工具调用 {tool_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]: return content[0]["text"] return 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