跳转至

mcp_base_tool.py — verl/tools/mcp_base_tool.py

文件路径

verl/tools/mcp_base_tool.py

文件概述

MCPBaseTool 是 MCP(Model Context Protocol)工具的基类,继承自 BaseTool。MCP 是一种标准化的协议,允许 LLM 通过统一接口与外部服务通信。相比于直接写代码调用 API 的 Native 工具,MCP 工具通过 MCP 客户端管理器(ClientManager)来代理所有工具调用,实现了更好的解耦和可扩展性。

通俗理解:如果说 BaseTool 是"自己动手写工具",那么 MCPBaseTool 就是"通过一个统一的中间人(MCP 协议)来调用远程工具服务"。

关键代码讲解

1. 构造函数

class MCPBaseTool(BaseTool):
    def __init__(self, config: dict, tool_schema: OpenAIFunctionToolSchema):
        super().__init__(config, tool_schema)
        self._instance_dict = {}
        self.timeout = config.get("timeout", 30)
  • 继承 BaseTool 的初始化逻辑。
  • _instance_dict:存储每个轨迹实例的状态(响应内容、奖励列表等)。
  • timeout:MCP 调用的超时时间,默认 30 秒。

2. 创建实例

async def create(self, instance_id: Optional[str] = None, **kwargs) -> tuple[str, ToolResponse]:
    if instance_id is None:
        instance_id = str(uuid4())
    self._instance_dict[instance_id] = {
        "response": "",
        "reward": [],
    }
    return instance_id, ToolResponse()

与 BaseTool 类似,但额外在实例字典中初始化了 response 和 reward 字段。

3. 调用 MCP 工具(核心方法)

async def _call_tool(self, instance_id, parameters) -> tuple[str, dict]:
    err_msg = ""
    metadata = {}
    try:
        call_tool_result = await ClientManager.call_tool(self.name, parameters, self.timeout)
        result, metadata = self._parse_tool_result(call_tool_result.content)
    except ClientError as e:
        err_msg = f"\n Tool call failed: {e}"
    except ConnectionError as e:
        err_msg = f"\n Connection failed: {e}"
    except Exception as e:
        err_msg = f"\n An unexpected error occurred: {e}"
    finally:
        if err_msg:
            result = err_msg
            metadata["api_request_error"] = err_msg
        else:
            metadata["api_request_error"] = None
    return result, metadata

关键步骤: 1. 通过全局的 ClientManager 单例调用 MCP 远程工具。 2. 解析返回结果(_parse_tool_result)。 3. 统一处理各种异常(客户端错误、连接错误等)。

4. 执行工具

@rollout_trace_op
async def execute(self, instance_id: str, parameters: dict[str, Any], **kwargs) -> tuple[ToolResponse, float, dict]:
    if self.name == "" or self.name is None or parameters is None:
        error_msg = "Error: 'parameters' is missing or empty."
        return ToolResponse(text=json.dumps({"result": error_msg})), 0.0, {}

    try:
        result_text, metadata = await self._call_tool(instance_id, parameters)
        self._instance_dict[instance_id]["reward"].append(result_text.strip())
        metrics = {
            "query_count": metadata.get("query_count", 0),
            "status": metadata.get("status", "unknown"),
            "total_results": metadata.get("total_results", 0),
            "api_request_error": metadata.get("api_request_error"),
        }
        return ToolResponse(text=result_text), 0.0, metrics
    except Exception as e:
        error_result = json.dumps({"result": f"Tool execution failed: {e}"})
        return ToolResponse(text=error_result), 0.0, {"error": str(e)}

先验证参数,然后调用 _call_tool 执行远程工具,将结果存入实例字典,并返回响应、奖励分数和度量指标。

5. 解析工具结果

def _parse_tool_result(self, content: list) -> tuple[str, dict]:
    tools_content = [part.text for part in filter(lambda x: x.type == "text", content)]
    return " ".join(tools_content), {}

从 MCP 返回的 content 列表中提取所有文本类型的内容,拼接成字符串。子类(如 MCPSearchTool)可以覆写此方法来实现自定义解析逻辑。

核心类/函数列表

类/方法 作用
MCPBaseTool MCP 工具的基类
_call_tool 通过 ClientManager 调用远程 MCP 工具
execute 执行工具(含参数验证和错误处理)
_parse_tool_result 解析 MCP 返回结果(可被子类覆写)
calc_reward 返回累积的奖励列表
release 释放实例并清理字典

与其他模块的关系

  • 继承自 BaseTool(base_tool.py)。
  • 被 MCPSearchTool(mcp_search_tool.py)继承。
  • 依赖 ClientManager(utils/mcp_clients/McpClientManager.py)进行远程调用。
  • 依赖 schemas.py 中的 OpenAIFunctionToolSchema 和 ToolResponse。

小结

MCPBaseTool 是通过 MCP 协议调用远程工具的基类。它封装了 MCP 客户端调用、错误处理和结果解析的通用逻辑,使得具体的 MCP 工具(如搜索工具)只需关注结果解析的定制化部分。