openenv/coding_env
21
1# Copyright (c) Meta Platforms, Inc. and affiliates.2# All rights reserved.3#4# This source code is licensed under the BSD-style license found in the5# LICENSE file in the root directory of this source tree.6 7"""8MCP Client classes for tool-calling environments.9 10This module provides async client classes for interacting with MCP-enabled environments:11- MCPClientBase: Base class with shared tool discovery12- MCPToolClient: Client for tool-calling style (one tool per step)13 14These clients abstract away the MCP protocol details, providing a clean interface15for listing and calling tools on remote environments. All clients are async by default.16 17Architecture Overview::18 19 ┌─────────────────────────────────────────────────────────┐20 │ HTTPEnvServer │21 ├─────────────────────────────────────────────────────────┤22 │ Simulation Mode (default): │23 │ /ws → OpenEnv protocol (reset/step/state) │24 │ /mcp → MCP JSON-RPC (tools/list, tools/call) │25 │ /reset, /step, /state → HTTP endpoints │26 ├─────────────────────────────────────────────────────────┤27 │ Production Mode (use_production_mode=True): │28 │ /mcp → MCP JSON-RPC (tools/list, tools/call) │29 │ Bypasses step() for direct tool access │30 └─────────────────────────────────────────────────────────┘31 32 Client Usage:33 MCPToolClient (default) → /ws (step-based, with rewards)34 MCPToolClient (production) → /mcp (direct tool access, no rewards)35 36Example (async):37 >>> from openenv.core.mcp_client import MCPToolClient38 >>>39 >>> async with MCPToolClient(base_url="http://localhost:8000") as env:40 ... # Discover available tools41 ... tools = await env.list_tools()42 ... print([t.name for t in tools])43 ...44 ... # Call a tool45 ... result = await env.call_tool("echo_message", message="Hello!")46 ... print(result)47 48Example (sync wrapper):49 >>> env = MCPToolClient(base_url="http://localhost:8000").sync()50 >>> with env:51 ... tools = env.list_tools()52 ... result = env.call_tool("echo_message", message="Hello!")53"""54 55import asyncio56from typing import Any, Dict, List, Optional57 58from .client_types import StepResult59from .env_client import EnvClient60from .env_server.mcp_types import (61 CallToolAction,62 CallToolObservation,63 ListToolsAction,64 ListToolsObservation,65 Tool,66 ToolError,67)68from .env_server.types import Observation, State69 70 71class MCPClientBase(EnvClient[Any, Observation, State]):72 """73 Base class for MCP clients with tool discovery.74 75 This class provides the common `list_tools()` method for discovering76 available tools from an MCP-enabled environment. Subclasses implement77 specific interaction patterns (tool-calling or CodeAct).78 79 Attributes:80 _tools_cache: Cached list of tools (populated on first `list_tools()` call)81 """82 83 def __init__(84 self,85 base_url: str,86 connect_timeout_s: float = 10.0,87 message_timeout_s: float = 60.0,88 provider: Optional[Any] = None,89 mode: Optional[str] = None,90 ):91 """92 Initialize MCP client.93 94 Args:95 base_url: Base URL of the environment server (http:// or ws://).96 connect_timeout_s: Timeout for establishing WebSocket connection.97 message_timeout_s: Timeout for receiving responses to messages.98 provider: Optional container/runtime provider for lifecycle management.99 mode: Communication mode. Must be 'production' for MCP clients. Defaults to 'production'.100 """101 # MCPClientBase defaults to production mode, but allow override for validation102 if mode is None:103 mode = "production"104 105 # Validate that mode is production106 mode_lower = mode.lower()107 if mode_lower != "production":108 raise ValueError(109 f"MCPToolClient only supports 'production' mode, got '{mode}'. "110 f"Use GenericEnvClient for simulation mode."111 )112 113 super().__init__(114 base_url=base_url,115 connect_timeout_s=connect_timeout_s,116 message_timeout_s=message_timeout_s,117 provider=provider,118 mode=mode,119 )120 self._tools_cache: Optional[List[Tool]] = None121 self.use_production_mode = False122 self._production_session_id: Optional[str] = None123 self._production_session_lock = asyncio.Lock()124 self._jsonrpc_request_id = 0125 self._http_client: Optional[Any] = None # lazily-created httpx.AsyncClient126 127 def _next_request_id(self) -> int:128 """Generate a monotonically increasing JSON-RPC request id."""129 self._jsonrpc_request_id += 1130 return self._jsonrpc_request_id131 132 def _production_mcp_url(self) -> str:133 """Build HTTP MCP endpoint URL from the client's websocket URL."""134 url = self._ws_url.replace("ws://", "http://").replace("wss://", "https://")135 if url.endswith("/ws"):136 url = url[: -len("/ws")]137 return url.rstrip("/") + "/mcp"138 139 async def _get_http_client(self) -> Any:140 """Return a shared httpx.AsyncClient, creating one lazily."""141 if self._http_client is None:142 import httpx143 144 self._http_client = httpx.AsyncClient()145 return self._http_client146 147 async def _production_mcp_request(148 self, method: str, params: Optional[Dict[str, Any]] = None149 ) -> Dict[str, Any]:150 """Send a JSON-RPC request to HTTP /mcp and return parsed JSON response."""151 client = await self._get_http_client()152 response = await client.post(153 self._production_mcp_url(),154 json={155 "jsonrpc": "2.0",156 "method": method,157 "params": params or {},158 "id": self._next_request_id(),159 },160 timeout=self._message_timeout,161 )162 response.raise_for_status()163 return response.json()164 165 async def _ensure_production_session(self) -> str:166 """Create and cache a persistent HTTP MCP session id if needed."""167 async with self._production_session_lock:168 if self._production_session_id is not None:169 return self._production_session_id170 171 data = await self._production_mcp_request("openenv/session/create")172 if "error" in data:173 message = data.get("error", {}).get("message", "unknown error")174 raise RuntimeError(f"Failed to create MCP session: {message}")175 176 session_id = data.get("result", {}).get("session_id")177 if not session_id:178 raise RuntimeError("Failed to create MCP session: missing session_id")179 180 self._production_session_id = session_id181 return session_id182 183 async def list_tools(self, use_cache: bool = True) -> List[Tool]:184 """185 Discover available tools from the environment.186 187 Args:188 use_cache: If True, return cached tools if available.189 Set to False to force a fresh request.190 191 Returns:192 List of Tool objects with name, description, and input_schema.193 194 Example:195 >>> tools = await env.list_tools()196 >>> for tool in tools:197 ... print(f"{tool.name}: {tool.description}")198 """199 if use_cache and self._tools_cache is not None:200 return self._tools_cache201 202 # Use production mode HTTP endpoint if enabled.203 # Some tests instantiate with __new__ and skip __init__, so default missing flag to False.204 if getattr(self, "use_production_mode", False):205 try:206 session_id = await self._ensure_production_session()207 data = await self._production_mcp_request(208 "tools/list",209 {"session_id": session_id},210 )211 if "error" in data:212 message = data.get("error", {}).get("message", "unknown error")213 raise RuntimeError(f"list_tools failed: {message}")214 if "result" in data and "tools" in data["result"]:215 tools = [216 Tool(217 name=t.get("name", ""),218 description=t.get("description", ""),219 input_schema=t.get(220 "input_schema", t.get("inputSchema", {})221 ),222 )223 for t in data["result"]["tools"]224 ]225 self._tools_cache = tools226 return tools227 except Exception:228 # If HTTP request fails, return empty list229 pass230 return []231 232 result = await self.step(ListToolsAction())233 if isinstance(result.observation, ListToolsObservation):234 self._tools_cache = result.observation.tools235 return self._tools_cache236 237 # Unexpected observation type; keep API stable with an empty tool list.238 self._tools_cache = []239 return self._tools_cache240 241 def _step_payload(self, action: Any) -> Dict[str, Any]:242 """Convert an Action object to the JSON data expected by the env server."""243 if isinstance(action, ListToolsAction):244 return {"type": "list_tools"}245 elif isinstance(action, CallToolAction):246 return {247 "type": "call_tool",248 "tool_name": action.tool_name,249 "arguments": action.arguments,250 }251 else:252 # For unknown actions, try to serialize as dict253 if hasattr(action, "model_dump"):254 return action.model_dump()255 return {"action": str(action)}256 257 def _parse_result(self, payload: Dict[str, Any]) -> StepResult[Observation]:258 """Convert a JSON response from the env server to StepResult[Observation]."""259 obs_data = payload.get("observation", {})260 261 # Check if this is a ListToolsObservation262 if "tools" in obs_data:263 tools = [264 Tool(265 name=t.get("name", ""),266 description=t.get("description", ""),267 input_schema=t.get("input_schema", t.get("inputSchema", {})),268 )269 for t in obs_data.get("tools", [])270 ]271 observation = ListToolsObservation(272 tools=tools,273 done=payload.get("done", False),274 reward=payload.get("reward"),275 metadata=obs_data.get("metadata", {}),276 )277 # Check if this is a CallToolObservation278 elif "tool_name" in obs_data:279 error = None280 if obs_data.get("error"):281 error = ToolError(**obs_data["error"])282 283 observation = CallToolObservation(284 tool_name=obs_data.get("tool_name", ""),285 result=obs_data.get("result"),286 error=error,287 done=payload.get("done", False),288 reward=payload.get("reward"),289 metadata=obs_data.get("metadata", {}),290 )291 else:292 # Generic observation293 observation = Observation(294 done=payload.get("done", False),295 reward=payload.get("reward"),296 metadata=obs_data.get("metadata", {}),297 )298 299 return StepResult(300 observation=observation,301 reward=payload.get("reward"),302 done=payload.get("done", False),303 )304 305 def _parse_state(self, payload: Dict[str, Any]) -> State:306 """Convert a JSON response from the state endpoint to a State object."""307 return State(308 episode_id=payload.get("episode_id"),309 step_count=payload.get("step_count", 0),310 )311 312 async def close(self) -> None:313 """314 Close client resources.315 316 In production MCP mode, this also closes the server-side persistent317 MCP session (best effort) before closing websocket/provider resources.318 """319 if self._production_session_id is not None:320 try:321 await self._production_mcp_request(322 "openenv/session/close",323 {"session_id": self._production_session_id},324 )325 except Exception:326 # Best effort cleanup - do not mask normal close behavior327 pass328 finally:329 self._production_session_id = None330 331 if self._http_client is not None:332 try:333 await self._http_client.aclose()334 except Exception:335 pass336 finally:337 self._http_client = None338 339 await super().close()340 341 342class MCPToolClient(MCPClientBase):343 """344 Async client for tool-calling style MCP interactions.345 346 Each step invokes a single tool. Use this for traditional function-calling347 agent patterns where the agent decides which tool to call next.348 349 This client provides convenience methods for tool discovery and invocation:350 - `list_tools()`: Get all available tools with their schemas351 - `call_tool(name, **kwargs)`: Invoke a tool by name with arguments352 353 Example (async):354 >>> async with MCPToolClient(base_url="http://localhost:8000") as env:355 ... # Reset the environment356 ... await env.reset()357 ...358 ... # Discover available tools359 ... tools = await env.list_tools()360 ... print([t.name for t in tools]) # ['echo_message', 'echo_with_length']361 ...362 ... # Call a tool directly363 ... result = await env.call_tool("echo_message", message="Hello!")364 ... print(result) # "Hello!"365 ...366 ... # Or use the full action interface367 ... from openenv.core.env_server.mcp_types import CallToolAction368 ... step_result = await env.step(CallToolAction(369 ... tool_name="echo_with_length",370 ... arguments={"message": "Test"}371 ... ))372 ... print(step_result.observation.result)373 374 Example (sync wrapper):375 >>> env = MCPToolClient(base_url="http://localhost:8000").sync()376 >>> with env:377 ... tools = env.list_tools()378 ... result = env.call_tool("echo_message", message="Hello!")379 """380 381 async def call_tool(self, name: str, **kwargs: Any) -> Any:382 """383 Call a tool by name.384 385 This is a convenience method that creates a CallToolAction, executes it,386 and returns the result directly. For more control, use `step()` with387 a CallToolAction directly.388 389 Args:390 name: Name of the tool to invoke (must match a tool from `list_tools()`).391 **kwargs: Arguments to pass to the tool. Must match the tool's input_schema.392 393 Returns:394 The tool's result. The type depends on the tool being called.395 396 Raises:397 RuntimeError: If the server returns an error response.398 399 Example:400 >>> result = await env.call_tool("add", a=5, b=3)401 >>> print(result) # 8402 >>>403 >>> result = await env.call_tool("greet", name="Claude")404 >>> print(result) # "Hello, Claude!"405 """406 if getattr(self, "use_production_mode", False):407 session_id = await self._ensure_production_session()408 data = await self._production_mcp_request(409 "tools/call",410 {411 "name": name,412 "arguments": kwargs,413 "session_id": session_id,414 },415 )416 417 if "error" in data:418 message = data.get("error", {}).get("message", "unknown error")419 raise RuntimeError(f"Tool '{name}' failed: {message}")420 421 result = data.get("result")422 if isinstance(result, dict) and "data" in result:423 return result["data"]424 return result425 426 action = CallToolAction(tool_name=name, arguments=kwargs)427 result = await self.step(action)428 obs = result.observation429 430 # Check for transport/framework errors431 if isinstance(obs, CallToolObservation) and obs.error is not None:432 raise RuntimeError(433 f"Tool '{name}' failed: {obs.error.message} "434 f"(type: {obs.error.error_type.value})"435 )436 437 # Return the result438 if isinstance(obs, CallToolObservation):439 result = obs.result440 # Handle FastMCP CallToolResult objects441 # - As object: has .data attribute442 # - As dict (from JSON): has "data" key443 if hasattr(result, "data"):444 return result.data445 if isinstance(result, dict) and "data" in result:446 return result["data"]447 return result448 449 # Fallback for unexpected observation types450 return obs451 452 async def get_tool(self, name: str) -> Optional[Tool]:453 """454 Get a specific tool by name.455 456 Args:457 name: Name of the tool to find.458 459 Returns:460 The Tool object if found, None otherwise.461 462 Example:463 >>> tool = await env.get_tool("echo_message")464 >>> if tool:465 ... print(tool.description)466 ... print(tool.input_schema)467 """468 tools = await self.list_tools()469 for tool in tools:470 if tool.name == name:471 return tool472 return None473 474 async def has_tool(self, name: str) -> bool:475 """476 Check if a tool exists.477 478 Args:479 name: Name of the tool to check.480 481 Returns:482 True if the tool exists, False otherwise.483 """484 return await self.get_tool(name) is not None485 