openenv/echo_env
6
1# SPDX-License-Identifier: BSD-3-Clause2 3"""4MCP Client classes for tool-calling environments.5 6This module provides async client classes for interacting with MCP-enabled environments:7- MCPClientBase: Base class with shared tool discovery8- MCPToolClient: Client for tool-calling style (one tool per step)9 10These clients abstract away the MCP protocol details, providing a clean interface11for listing and calling tools on remote environments. All clients are async by default.12 13Architecture Overview::14 15 ┌─────────────────────────────────────────────────────────┐16 │ HTTPEnvServer │17 ├─────────────────────────────────────────────────────────┤18 │ Simulation Mode (default): │19 │ /ws → OpenEnv protocol (reset/step/state) │20 │ /mcp → MCP JSON-RPC (tools/list, tools/call) │21 │ /reset, /step, /state → HTTP endpoints │22 ├─────────────────────────────────────────────────────────┤23 │ Production Mode (use_production_mode=True): │24 │ /mcp → MCP JSON-RPC (tools/list, tools/call) │25 │ Bypasses step() for direct tool access │26 └─────────────────────────────────────────────────────────┘27 28 Client Usage:29 MCPToolClient (default) → /ws (step-based, with rewards)30 MCPToolClient (production) → /mcp (direct tool access, no rewards)31 32Examples:33 34 ```python35 from openenv.core.mcp_client import MCPToolClient36 37 async with MCPToolClient(base_url="http://localhost:8000") as env:38 # Discover available tools39 tools = await env.list_tools()40 print([t.name for t in tools])41 42 # Call a tool43 result = await env.call_tool("echo_message", message="Hello!")44 print(result)45 ```46 47 Sync wrapper:48 49 ```python50 env = MCPToolClient(base_url="http://localhost:8000").sync()51 with env:52 tools = env.list_tools()53 result = env.call_tool("echo_message", message="Hello!")54 ```55"""56 57import asyncio58from typing import Any, Dict, List, Optional59from urllib.parse import parse_qsl, urlencode, urlsplit, urlunsplit60 61from pydantic import ConfigDict62 63from .client_types import StepResult64from .env_client import EnvClient65from .env_server.mcp_types import (66 CallToolAction,67 CallToolObservation,68 ListToolsAction,69 ListToolsObservation,70 Tool,71 ToolError,72)73from .env_server.types import Observation, State74 75 76def _tool_from_payload(payload: Dict[str, Any]) -> Tool:77 """Convert a JSON tool payload into the internal Tool model."""78 return Tool(79 name=payload.get("name", ""),80 description=payload.get("description", ""),81 input_schema=payload.get("input_schema", payload.get("inputSchema", {})),82 )83 84 85class GenericMCPObservation(Observation):86 """Generic MCP observation that preserves env-specific fields."""87 88 model_config = ConfigDict(89 extra="allow",90 validate_assignment=True,91 arbitrary_types_allowed=True,92 )93 94 95class MCPClientBase(EnvClient[Any, Observation, State]):96 """97 Base class for MCP clients with tool discovery.98 99 This class provides the common `list_tools()` method for discovering100 available tools from an MCP-enabled environment. Subclasses implement101 specific interaction patterns (tool-calling or CodeAct).102 103 Attributes:104 _tools_cache: Cached list of tools (populated on first `list_tools()` call)105 """106 107 def __init__(108 self,109 base_url: str,110 connect_timeout_s: float = 10.0,111 message_timeout_s: float = 60.0,112 websocket_ping_interval_s: Optional[float] = 20.0,113 websocket_ping_timeout_s: Optional[float] = 20.0,114 provider: Optional[Any] = None,115 mode: Optional[str] = None,116 max_message_size_mb: float = 100.0,117 ):118 """119 Initialize MCP client.120 121 Args:122 base_url (`str`):123 Base URL of the environment server (http:// or ws://).124 connect_timeout_s (`float`, *optional*, defaults to `10.0`):125 Timeout for establishing WebSocket connection.126 message_timeout_s (`float`, *optional*, defaults to `60.0`):127 Timeout for receiving responses to messages.128 websocket_ping_interval_s (`float` or `None`, *optional*, defaults to `20.0`):129 WebSocket keepalive ping interval. Pass `None` to disable.130 websocket_ping_timeout_s (`float` or `None`, *optional*, defaults to `20.0`):131 WebSocket keepalive pong timeout. Pass `None` to disable.132 provider (*optional*):133 Container/runtime provider for lifecycle management.134 mode (`str`, *optional*):135 Communication mode. Must be 'production' for MCP clients. Defaults to 'production'.136 max_message_size_mb (`float`, *optional*, defaults to `100.0`):137 Largest WebSocket frame to accept. `EnvClient` has always taken this, but138 `MCPClientBase` did not forward it, so no MCP client could raise it — an environment139 whose tool returns a large result closed the connection with `1009 message too big`140 and there was no way to ask for more from the client side.141 """142 # MCPClientBase defaults to production mode, but allow override for validation143 if mode is None:144 mode = "production"145 146 # Validate that mode is production147 mode_lower = mode.lower()148 if mode_lower != "production":149 raise ValueError(150 f"MCPToolClient only supports 'production' mode, got '{mode}'. "151 f"Use GenericEnvClient for simulation mode."152 )153 154 super().__init__(155 base_url=base_url,156 connect_timeout_s=connect_timeout_s,157 message_timeout_s=message_timeout_s,158 websocket_ping_interval_s=websocket_ping_interval_s,159 websocket_ping_timeout_s=websocket_ping_timeout_s,160 provider=provider,161 mode=mode,162 max_message_size_mb=max_message_size_mb,163 )164 self._tools_cache: Optional[List[Tool]] = None165 self.use_production_mode = self._mode == "production"166 self._production_session_id: Optional[str] = None167 self._production_connect_lock = asyncio.Lock()168 self._production_session_lock = asyncio.Lock()169 self._jsonrpc_request_id = 0170 self._http_client: Optional[Any] = None # lazily-created httpx.AsyncClient171 172 def _next_request_id(self) -> int:173 """Generate a monotonically increasing JSON-RPC request id."""174 self._jsonrpc_request_id += 1175 return self._jsonrpc_request_id176 177 def _production_mcp_url(self) -> str:178 """Build the HTTP MCP endpoint URL from the stable base URL."""179 if self._base_url is None:180 raise RuntimeError("MCP client is not connected to a server.")181 parts = urlsplit(self._base_url)182 scheme = {"ws": "http", "wss": "https"}.get(parts.scheme, parts.scheme)183 return urlunsplit(184 parts._replace(185 scheme=scheme,186 path=parts.path.rstrip("/") + "/mcp",187 query="",188 fragment="",189 )190 )191 192 async def _get_http_client(self) -> Any:193 """Return a shared httpx.AsyncClient, creating one lazily."""194 if self._http_client is None:195 import httpx196 197 self._http_client = httpx.AsyncClient()198 return self._http_client199 200 async def _production_mcp_request(201 self, method: str, params: Optional[Dict[str, Any]] = None202 ) -> Dict[str, Any]:203 """Send a JSON-RPC request to HTTP /mcp and return parsed JSON response."""204 client = await self._get_http_client()205 response = await client.post(206 self._production_mcp_url(),207 json={208 "jsonrpc": "2.0",209 "method": method,210 "params": params or {},211 "id": self._next_request_id(),212 },213 timeout=self._message_timeout,214 )215 response.raise_for_status()216 return response.json()217 218 async def _connect_async(self) -> EnvClient:219 """220 Establish connection to the server.221 222 In production mode (use_production_mode=True), creates an HTTP MCP session223 and connects the WebSocket using that session ID so that WebSocket (reset/step/state)224 and HTTP MCP (list_tools/call_tool) share the exact same server-side environment session.225 """226 if getattr(self, "use_production_mode", False):227 async with self._production_connect_lock:228 try:229 self._start_provider_if_needed()230 session_id = await self._ensure_production_session()231 original_ws_url = self._ws_url232 if original_ws_url is None:233 raise RuntimeError("MCP client has no WebSocket URL.")234 235 parts = urlsplit(original_ws_url)236 query = dict(parse_qsl(parts.query, keep_blank_values=True))237 query["session_id"] = session_id238 self._ws_url = urlunsplit(parts._replace(query=urlencode(query)))239 try:240 await super()._connect_async()241 finally:242 self._ws_url = original_ws_url243 except BaseException:244 # Cancellation after the HTTP session is allocated must245 # release that session and any started provider before the246 # cancellation propagates.247 await self.close()248 raise249 return self250 251 return await super()._connect_async()252 253 async def _ensure_production_session(self) -> str:254 """Create and cache a persistent HTTP MCP session id if needed."""255 async with self._production_session_lock:256 if self._production_session_id is not None:257 return self._production_session_id258 259 data = await self._production_mcp_request("openenv/session/create")260 if "error" in data:261 message = data.get("error", {}).get("message", "unknown error")262 raise RuntimeError(f"Failed to create MCP session: {message}")263 264 session_id = data.get("result", {}).get("session_id")265 if not session_id:266 raise RuntimeError("Failed to create MCP session: missing session_id")267 268 self._production_session_id = session_id269 return session_id270 271 async def list_tools(self, use_cache: bool = True) -> List[Tool]:272 """273 Discover available tools from the environment.274 275 Args:276 use_cache (`bool`, *optional*, defaults to `True`):277 If `True`, return cached tools if available. Set to `False` to force a fresh request.278 279 Returns:280 List of `Tool` objects with name, description, and input_schema.281 282 Examples:283 284 ```python285 tools = await env.list_tools()286 for tool in tools:287 print(f"{tool.name}: {tool.description}")288 ```289 """290 if use_cache and self._tools_cache is not None:291 return self._tools_cache292 293 # Use production mode HTTP endpoint if enabled.294 # Some tests instantiate with __new__ and skip __init__, so default missing flag to False.295 if getattr(self, "use_production_mode", False):296 try:297 session_id = await self._ensure_production_session()298 data = await self._production_mcp_request(299 "tools/list",300 {"session_id": session_id},301 )302 if "error" in data:303 message = data.get("error", {}).get("message", "unknown error")304 raise RuntimeError(f"list_tools failed: {message}")305 if "result" in data and "tools" in data["result"]:306 tools = [_tool_from_payload(t) for t in data["result"]["tools"]]307 self._tools_cache = tools308 return tools309 except Exception:310 # If HTTP request fails, return empty list311 pass312 return []313 314 result = await self.step(ListToolsAction())315 if isinstance(result.observation, ListToolsObservation):316 self._tools_cache = result.observation.tools317 return self._tools_cache318 319 # Unexpected observation type; keep API stable with an empty tool list.320 self._tools_cache = []321 return self._tools_cache322 323 def _step_payload(self, action: Any) -> Dict[str, Any]:324 """Convert an Action object to the JSON data expected by the env server."""325 if isinstance(action, ListToolsAction):326 return {"type": "list_tools"}327 elif isinstance(action, CallToolAction):328 return {329 "type": "call_tool",330 "tool_name": action.tool_name,331 "arguments": action.arguments,332 }333 else:334 # For unknown actions, try to serialize as dict335 if hasattr(action, "model_dump"):336 return action.model_dump()337 return {"action": str(action)}338 339 def _parse_result(self, payload: Dict[str, Any]) -> StepResult[Observation]:340 """Convert a JSON response from the env server to StepResult[Observation]."""341 obs_data = payload.get("observation", {})342 343 # Check if this is a ListToolsObservation344 if "tools" in obs_data:345 tools = [_tool_from_payload(t) for t in obs_data.get("tools", [])]346 observation = ListToolsObservation(347 tools=tools,348 done=payload.get("done", False),349 reward=payload.get("reward"),350 metadata=payload.get("metadata", obs_data.get("metadata", {})),351 )352 # Check if this is a CallToolObservation353 elif "tool_name" in obs_data:354 error = None355 if obs_data.get("error"):356 error = ToolError(**obs_data["error"])357 358 observation = CallToolObservation(359 tool_name=obs_data.get("tool_name", ""),360 result=obs_data.get("result"),361 error=error,362 done=payload.get("done", False),363 reward=payload.get("reward"),364 metadata=payload.get("metadata", obs_data.get("metadata", {})),365 )366 else:367 # Generic observation with passthrough for env-specific fields368 custom_fields = {369 key: value370 for key, value in obs_data.items()371 if key not in {"done", "reward", "metadata"}372 }373 observation = GenericMCPObservation(374 done=payload.get("done", False),375 reward=payload.get("reward"),376 metadata=payload.get("metadata", obs_data.get("metadata", {})),377 **custom_fields,378 )379 380 return StepResult(381 observation=observation,382 reward=payload.get("reward"),383 done=payload.get("done", False),384 metadata=payload.get("metadata"),385 )386 387 def _parse_state(self, payload: Dict[str, Any]) -> State:388 """Convert a JSON response from the state endpoint to a State object."""389 return State(390 episode_id=payload.get("episode_id"),391 step_count=payload.get("step_count", 0),392 )393 394 async def _close_async(self) -> None:395 """396 Close client resources.397 398 In production MCP mode, this also closes the server-side persistent399 MCP session (best effort) after detaching the WebSocket and before400 closing HTTP/provider resources.401 402 Override `_close_async` rather than `close` so sync teardown403 (`SyncEnvClient.close`, sync `__exit__`, and `_dispatch`) still cleans404 up the HTTP MCP session.405 """406 try:407 # The WebSocket shares the HTTP-created session. Detach it first so408 # the server's ownership guard permits the explicit session close.409 await self._disconnect_async()410 finally:411 try:412 if self._production_session_id is not None:413 try:414 await self._production_mcp_request(415 "openenv/session/close",416 {"session_id": self._production_session_id},417 )418 except Exception:419 # Best effort cleanup - do not mask normal close behavior420 pass421 finally:422 self._production_session_id = None423 finally:424 try:425 if self._http_client is not None:426 try:427 await self._http_client.aclose()428 except Exception:429 pass430 finally:431 self._http_client = None432 finally:433 # This is intentionally inside the outer finally so434 # cancellation cannot skip provider teardown.435 await super()._close_async()436 437 438class MCPToolClient(MCPClientBase):439 """440 Async client for tool-calling style MCP interactions.441 442 Each step invokes a single tool. Use this for traditional function-calling443 agent patterns where the agent decides which tool to call next.444 445 This client provides convenience methods for tool discovery and invocation:446 - `list_tools()`: Get all available tools with their schemas447 - `call_tool(name, **kwargs)`: Invoke a tool by name with arguments448 449 Examples:450 451 ```python452 async with MCPToolClient(base_url="http://localhost:8000") as env:453 # Reset the environment454 await env.reset()455 456 # Discover available tools457 tools = await env.list_tools()458 print([t.name for t in tools]) # ['echo_message', 'echo_with_length']459 460 # Call a tool directly461 result = await env.call_tool("echo_message", message="Hello!")462 print(result) # "Hello!"463 464 # Or use the full action interface465 from openenv.core.env_server.mcp_types import CallToolAction466 step_result = await env.step(CallToolAction(467 tool_name="echo_with_length",468 arguments={"message": "Test"}469 ))470 print(step_result.observation.result)471 ```472 473 Sync wrapper:474 475 ```python476 env = MCPToolClient(base_url="http://localhost:8000").sync()477 with env:478 tools = env.list_tools()479 result = env.call_tool("echo_message", message="Hello!")480 ```481 """482 483 async def call_tool(self, name: str, **kwargs: Any) -> Any:484 """485 Call a tool by name.486 487 This is a convenience method that creates a CallToolAction, executes it,488 and returns the result directly. For more control, use `step()` with489 a CallToolAction directly.490 491 Args:492 name (`str`):493 Name of the tool to invoke (must match a tool from `list_tools()`).494 **kwargs:495 Arguments to pass to the tool. Must match the tool's input_schema.496 497 Returns:498 The tool's result. The type depends on the tool being called.499 500 Raises:501 `RuntimeError`: If the server returns an error response.502 503 Examples:504 505 ```python506 result = await env.call_tool("add", a=5, b=3)507 print(result) # 8508 509 result = await env.call_tool("greet", name="Claude")510 print(result) # "Hello, Claude!"511 ```512 """513 if getattr(self, "use_production_mode", False):514 session_id = await self._ensure_production_session()515 data = await self._production_mcp_request(516 "tools/call",517 {518 "name": name,519 "arguments": kwargs,520 "session_id": session_id,521 },522 )523 524 if "error" in data:525 message = data.get("error", {}).get("message", "unknown error")526 raise RuntimeError(f"Tool '{name}' failed: {message}")527 528 result = data.get("result")529 if isinstance(result, dict) and "data" in result:530 return result["data"]531 return result532 533 action = CallToolAction(tool_name=name, arguments=kwargs)534 result = await self.step(action)535 obs = result.observation536 537 # Check for transport/framework errors538 if isinstance(obs, CallToolObservation) and obs.error is not None:539 raise RuntimeError(540 f"Tool '{name}' failed: {obs.error.message} "541 f"(type: {obs.error.error_type.value})"542 )543 544 # Return the result545 if isinstance(obs, CallToolObservation):546 result = obs.result547 # Handle FastMCP CallToolResult objects548 # - As object: has .data attribute549 # - As dict (from JSON): has "data" key550 if hasattr(result, "data"):551 return result.data552 if isinstance(result, dict) and "data" in result:553 return result["data"]554 return result555 556 # Fallback for unexpected observation types557 return obs558 559 async def get_tool(self, name: str) -> Optional[Tool]:560 """561 Get a specific tool by name.562 563 Args:564 name (`str`):565 Name of the tool to find.566 567 Returns:568 The `Tool` object if found, `None` otherwise.569 570 Examples:571 572 ```python573 tool = await env.get_tool("echo_message")574 if tool:575 print(tool.description)576 print(tool.input_schema)577 ```578 """579 tools = await self.list_tools()580 for tool in tools:581 if tool.name == name:582 return tool583 return None584 585 async def has_tool(self, name: str) -> bool:586 """587 Check if a tool exists.588 589 Args:590 name (`str`):591 Name of the tool to check.592 593 Returns:594 `True` if the tool exists, `False` otherwise.595 """596 return await self.get_tool(name) is not None597 