openenv/echo_env
6
1# SPDX-License-Identifier: BSD-3-Clause2 3"""4MCP Environment base class for OpenEnv.5 6This module provides the MCPEnvironment base class that integrates FastMCP servers7with OpenEnv's Gym-style Environment interface. It handles MCP tool discovery8and invocation through the step() API, following RFC 003.9 10Key features:11- Automatic routing of ListToolsAction and CallToolAction to MCP server12- Reserved tool name validation (reset, step, state, close are protected)13- Timeout handling for tool calls14- Proper error categorization (tool not found, execution errors, timeouts)15- Mode-aware tool registration (production vs simulation)16- Code mode support via get_callables() and execute_code()17 18Examples:19 20 ```python21 from fastmcp import FastMCP22 from openenv.core.env_server.mcp_environment import MCPEnvironment23 24 class MyMCPEnv(MCPEnvironment):25 def __init__(self):26 mcp = FastMCP("my-server")27 28 # Register mode-specific tools29 @self.tool(mode="production")30 def my_tool(arg: str) -> str:31 return f"Production: {arg}"32 33 @self.tool(mode="simulation")34 def my_tool(arg: str) -> str:35 return f"Simulation: {arg}"36 37 super().__init__(mcp)38 39 def reset(self, seed=None, episode_id=None, **kwargs):40 # Reset logic here41 ...42 43 def _step_impl(self, action):44 # Handle non-MCP actions45 ...46 47 @property48 def state(self):49 # Return current state50 ...51 ```52"""53 54import asyncio55import inspect56from abc import abstractmethod57from collections import defaultdict58from contextlib import asynccontextmanager59from typing import Any, Callable, Dict, Optional60 61from fastmcp import Client62from fastmcp.client.client import CallToolResult63from mcp.types import TextContent64 65from ..utils import run_async_safely66from .interfaces import Environment67from .mcp_types import (68 CallToolAction,69 CallToolObservation,70 ListToolsAction,71 ListToolsObservation,72 RESERVED_TOOL_NAMES,73 Tool,74 ToolError,75 ToolErrorType,76)77from .types import Action, Observation78 79 80# Default timeout for MCP tool calls in seconds81MCP_TOOL_CALL_TIMEOUT = 30.082 83# Valid modes for tool registration84VALID_MODES = {"production", "simulation"}85 86 87def get_server_tools(mcp_server: Any) -> Dict[str, Any]:88 """89 Get tools from a FastMCP server, compatible with both 2.x and 3.x.90 91 Returns:92 Dictionary mapping tool names to tool objects.93 """94 # FastMCP 2.x: get_tools() returns dict {name: Tool}95 if hasattr(mcp_server, "get_tools"):96 result = run_async_safely(mcp_server.get_tools())97 if isinstance(result, dict):98 return result99 # FastMCP 3.x: list_tools() returns list of Tool objects100 if hasattr(mcp_server, "list_tools"):101 tools_list = run_async_safely(mcp_server.list_tools())102 return {t.name: t for t in tools_list}103 return {}104 105 106def _tool_from_server_tool(tool: Any) -> Tool:107 """Convert a FastMCP tool object into OpenEnv's Tool model."""108 return Tool(109 name=tool.name,110 description=tool.description or "",111 input_schema=tool.inputSchema if hasattr(tool, "inputSchema") else {},112 )113 114 115def _tool_from_mode_schema(schema: Dict[str, Any]) -> Tool:116 """Convert a stored mode-aware schema into OpenEnv's Tool model."""117 return Tool(118 name=schema["name"],119 description=schema["description"],120 input_schema=schema["input_schema"],121 )122 123 124def _schema_for_mode(125 mode_schemas: Dict[Optional[str], Dict[str, Any]], current_mode: Optional[str]126) -> Optional[Dict[str, Any]]:127 """Return the schema visible for the current mode, if any."""128 if None in mode_schemas:129 return mode_schemas[None]130 return mode_schemas.get(current_mode)131 132 133class MCPEnvironment(Environment):134 """135 Base class for environments that expose tools via MCP (Model Context Protocol).136 137 MCPEnvironment bridges FastMCP servers with OpenEnv's Gym-style API, allowing138 agents to discover and invoke MCP tools through the standard step() interface.139 140 The class automatically handles:141 - ListToolsAction: Returns available tools from the MCP server142 - CallToolAction: Invokes a specific tool with arguments143 144 All other actions are delegated to the abstract _step_impl() method,145 which subclasses must implement.146 147 Args:148 mcp_server: A FastMCP server instance containing tool definitions.149 The server's tools will be validated against reserved names.150 transform: Optional transform to apply to observations (inherited from Environment).151 152 Raises:153 ValueError: If any tool in the MCP server uses a reserved name154 (reset, step, state, close).155 156 Examples:157 158 ```python159 from fastmcp import FastMCP160 161 mcp = FastMCP("calculator")162 163 @mcp.tool()164 def add(a: int, b: int) -> int:165 return a + b166 167 env = MyMCPEnvironment(mcp)168 obs = env.step(ListToolsAction())169 obs.tools[0].name # 'add'170 ```171 """172 173 def __init__(self, mcp_server: Any, transform: Optional[Any] = None) -> None:174 """175 Initialize the MCP environment.176 177 Args:178 mcp_server: A FastMCP server instance with tool definitions.179 transform: Optional transform to apply to observations.180 181 Raises:182 ValueError: If any tool uses a reserved name (reset, step, state, close).183 """184 super().__init__(transform=transform)185 186 # Validate tool names before storing187 self._validate_tool_names(mcp_server)188 189 self.mcp_server = mcp_server190 self.mcp_client = Client(mcp_server)191 192 # Track mode-specific tools: {tool_name: {mode: func}}193 # mode can be "production", "simulation", or None (available in all modes)194 self._mode_tools = defaultdict(dict)195 196 # Track tool schemas for list_tools: {tool_name: {mode: schema}}197 self._mode_tool_schemas = defaultdict(dict)198 199 def _require_mcp_client(self) -> Any:200 """Return MCP client or raise if environment has been closed."""201 if self.mcp_client is None:202 raise RuntimeError("MCP client is not available; environment is closed")203 return self.mcp_client204 205 def _require_mcp_server(self) -> Any:206 """Return MCP server or raise if environment has been closed."""207 if self.mcp_server is None:208 raise RuntimeError("MCP server is not available; environment is closed")209 return self.mcp_server210 211 @asynccontextmanager212 async def mcp_session(self):213 """214 Context manager for MCP client sessions.215 216 This wrapper serves two purposes:217 218 1. **Null guard** — raises a clear error if ``close()`` has already219 been called (``mcp_client`` is ``None``).220 221 2. **AsyncExitStack adapter** — FastMCP's ``Client.__aenter__``222 creates a background ``asyncio.Task`` for session management.223 When entered directly via ``AsyncExitStack`` in the HTTP session224 path (``_create_session``), this task can be cancelled by ASGI225 harnesses (e.g. Starlette ``TestClient``) between requests,226 corrupting session state. Wrapping in an ``asynccontextmanager``227 generator isolates the task lifecycle: the generator frame keeps228 ``async with client:`` suspended at ``yield``, so cleanup only229 runs when the stack explicitly closes the generator — not when230 the event loop cancels orphaned tasks.231 232 Delegates to FastMCP's ``Client`` context manager which is233 reentrant: the first entry opens the transport and subsequent234 (nested) entries simply increment an internal reference counter.235 The transport is closed only when the outermost context exits.236 237 No external lock is needed because ``Client._connect`` /238 ``Client._disconnect`` already serialise connection state changes239 through their own ``anyio.Lock``.240 """241 client = self._require_mcp_client()242 async with client:243 yield client244 245 @property246 def supports_code_mode(self) -> bool:247 """Check if this environment supports code mode (execute_code)."""248 return True249 250 def _get_server_tools(self, mcp_server: Any) -> Dict[str, Any]:251 """252 Get tools from a FastMCP server, compatible with both 2.x and 3.x.253 254 Returns:255 Dictionary mapping tool names to tool objects.256 """257 return get_server_tools(mcp_server)258 259 def get_callables(self) -> Dict[str, Callable]:260 """261 Get callable functions for code mode.262 263 Returns tool functions as direct Python callables, enabling code mode264 where agents write Python code that calls tools directly (no JSON-RPC265 overhead). Mode-specific tools are filtered by the current mode.266 267 Returns:268 Dictionary mapping tool names to callables.269 """270 callables: Dict[str, Callable] = {}271 current_mode = getattr(self, "_mode", None)272 273 # Extract callables from FastMCP server using public API274 for tool_name, tool in self._get_server_tools(self.mcp_server).items():275 if hasattr(tool, "fn") and callable(tool.fn):276 callables[tool_name] = tool.fn277 278 # Add mode-specific tools available in current mode279 for tool_name, mode_funcs in self._mode_tools.items():280 if None in mode_funcs:281 # Tool available in all modes (already in FastMCP if registered there)282 if tool_name not in callables:283 callables[tool_name] = mode_funcs[None]284 elif current_mode in mode_funcs:285 # Tool available in current mode only286 callables[tool_name] = mode_funcs[current_mode]287 288 return callables289 290 def execute_code(self, code: str) -> Observation:291 """292 Execute Python code with tools available as callables.293 294 This enables the CodeAct pattern where agents write Python code295 that calls tools directly as functions, avoiding JSON-RPC overhead.296 297 Args:298 code: Python code to execute. Tools are available as functions299 in the execution namespace. Set a variable named 'result'300 to capture the return value.301 302 Returns:303 Observation with result in metadata["result"] or error in304 metadata["error"].305 """306 namespace = self.get_callables()307 308 result_dict: Dict[str, Any] = {}309 try:310 exec(code, namespace, result_dict)311 result = result_dict.get("result")312 return Observation(done=False, reward=0.0, metadata={"result": result})313 except SyntaxError as e:314 return Observation(315 done=False, reward=0.0, metadata={"error": f"Syntax error: {str(e)}"}316 )317 except Exception as e:318 return Observation(done=False, reward=0.0, metadata={"error": str(e)})319 320 def _validate_tool_names(self, mcp_server: Any) -> None:321 """322 Validate that no tools use reserved names.323 324 Reserved names (reset, step, state, close) are protected to maintain325 the dual API boundary between infrastructure and agent APIs.326 327 Args:328 mcp_server: The FastMCP server to validate.329 330 Raises:331 ValueError: If any tool uses a reserved name.332 """333 tools_dict = self._get_server_tools(mcp_server)334 if tools_dict:335 tool_names = set(tools_dict.keys())336 conflicts = tool_names & RESERVED_TOOL_NAMES337 if conflicts:338 raise ValueError(339 f"MCP tools cannot use reserved names: {sorted(conflicts)}. "340 f"Reserved names are: {sorted(RESERVED_TOOL_NAMES)}"341 )342 343 def tool(self, mode: Optional[str] = None) -> Callable:344 """345 Decorator for registering mode-aware tools.346 347 Args:348 mode: Optional mode for the tool ("production" or "simulation").349 If None, tool is available in all modes.350 351 Returns:352 A decorator function for registering tools.353 354 Raises:355 ValueError: If mode is not None, "production", or "simulation".356 """357 if mode is not None and mode not in VALID_MODES:358 raise ValueError(359 f"Invalid mode '{mode}'. Mode must be 'production', 'simulation', or None."360 )361 362 def decorator(func: Callable) -> Callable:363 tool_name = func.__name__364 # Validate tool name is not reserved365 if tool_name in RESERVED_TOOL_NAMES:366 raise ValueError(367 f"Tool name '{tool_name}' is reserved and cannot be used. "368 f"Reserved names are: {sorted(RESERVED_TOOL_NAMES)}"369 )370 371 # If mode is None, register with FastMCP as usual372 if mode is None:373 mcp_server = self._require_mcp_server()374 decorated_func = mcp_server.tool()(func)375 self._mode_tools[tool_name][None] = func376 return decorated_func377 378 # For mode-specific tools, don't register with FastMCP379 # Instead, track them ourselves380 self._mode_tools[tool_name][mode] = func381 382 # Extract schema information from function signature383 sig = inspect.signature(func)384 schema = {385 "type": "object",386 "properties": {},387 "required": [],388 }389 390 for param_name, param in sig.parameters.items():391 # Get type annotation392 param_type = param.annotation393 json_type = "string" # default394 if param_type in (int, "int"):395 json_type = "integer"396 elif param_type in (float, "float"):397 json_type = "number"398 elif param_type in (bool, "bool"):399 json_type = "boolean"400 401 schema["properties"][param_name] = {"type": json_type}402 403 # If no default value, it's required404 if param.default == inspect.Parameter.empty:405 schema["required"].append(param_name)406 407 # Store the schema for this mode-specific tool408 self._mode_tool_schemas[tool_name][mode] = {409 "name": tool_name,410 "description": func.__doc__ or "",411 "input_schema": schema,412 }413 414 return func415 416 return decorator417 418 def step(419 self,420 action: Action,421 timeout_s: Optional[float] = None,422 **kwargs: Any,423 ) -> Observation:424 """425 Execute an action in the environment.426 427 This method routes MCP-specific actions (ListToolsAction, CallToolAction)428 to the appropriate handlers, while delegating all other actions to429 the subclass's _step_impl() method.430 431 Args:432 action (`Action`):433 The action to execute. `ListToolsAction` returns available MCP tools,434 `CallToolAction` invokes a specific MCP tool, and any other action435 is delegated to _step_impl().436 timeout_s (`float`, *optional*):437 Timeout in seconds for the action. Defaults to MCP_TOOL_CALL_TIMEOUT438 (30s) for MCP actions.439 **kwargs (`Any`):440 Additional arguments passed to handlers.441 442 Returns:443 `Observation`: `ListToolsObservation` for `ListToolsAction`,444 `CallToolObservation` for `CallToolAction`, or a subclass-defined445 Observation for other actions.446 """447 if isinstance(action, ListToolsAction):448 return self._handle_list_tools()449 elif isinstance(action, CallToolAction):450 return self._handle_call_tool(action, timeout_s=timeout_s)451 else:452 return self._step_impl(action, timeout_s=timeout_s, **kwargs)453 454 def _handle_list_tools(self) -> ListToolsObservation:455 """Sync wrapper — delegates to the canonical async implementation."""456 return run_async_safely(self._async_handle_list_tools())457 458 async def _async_list_tools(self) -> list:459 """460 Async helper to list tools from the MCP client.461 462 Returns:463 List of tool objects from the MCP server.464 """465 async with self.mcp_session() as client:466 return await client.list_tools()467 468 def _handle_call_tool(469 self,470 action: CallToolAction,471 timeout_s: Optional[float] = None,472 ) -> CallToolObservation:473 """Sync wrapper — delegates to the canonical async implementation."""474 return run_async_safely(475 self._async_handle_call_tool(action, timeout_s=timeout_s)476 )477 478 async def _async_call_tool(self, tool_name: str, arguments: dict) -> Any:479 """480 Async helper to call a tool on the MCP server.481 482 Args:483 tool_name: Name of the tool to invoke.484 arguments: Dictionary of arguments to pass to the tool.485 486 Returns:487 The result from the tool execution.488 """489 async with self.mcp_session() as client:490 return await client.call_tool(tool_name, arguments)491 492 async def _async_handle_list_tools(self) -> ListToolsObservation:493 """Async version of _handle_list_tools — avoids run_async_safely."""494 try:495 current_mode = getattr(self, "_mode", None)496 tools_result = await self._async_list_tools()497 tools = []498 for tool in tools_result:499 if tool.name not in self._mode_tool_schemas:500 tools.append(_tool_from_server_tool(tool))501 for mode_schemas in self._mode_tool_schemas.values():502 schema = _schema_for_mode(mode_schemas, current_mode)503 if schema is not None:504 tools.append(_tool_from_mode_schema(schema))505 return ListToolsObservation(tools=tools)506 except Exception as e:507 return ListToolsObservation(508 tools=[],509 metadata={"error": str(e), "error_type": "list_tools_failed"},510 )511 512 async def _async_handle_call_tool(513 self,514 action: CallToolAction,515 timeout_s: Optional[float] = None,516 ) -> CallToolObservation:517 """Async version of _handle_call_tool — avoids run_async_safely."""518 timeout = timeout_s if timeout_s is not None else MCP_TOOL_CALL_TIMEOUT519 tool_name = action.tool_name520 current_mode = getattr(self, "_mode", None)521 522 if tool_name in self._mode_tools:523 mode_info = self._mode_tools[tool_name]524 if None in mode_info:525 func = mode_info[None]526 elif current_mode in mode_info:527 func = mode_info[current_mode]528 else:529 return CallToolObservation(530 tool_name=tool_name,531 result=None,532 error=ToolError(533 error_type=ToolErrorType.TOOL_NOT_FOUND,534 message=f"Tool '{tool_name}' not available in {current_mode} mode",535 ),536 )537 try:538 if inspect.iscoroutinefunction(func):539 result = await func(**action.arguments)540 else:541 result = func(**action.arguments)542 return CallToolObservation(543 tool_name=tool_name,544 result=CallToolResult(545 content=[TextContent(type="text", text=str(result))],546 structured_content={"result": result},547 meta=None,548 data=result,549 is_error=False,550 ),551 )552 except Exception as e:553 return CallToolObservation(554 tool_name=tool_name,555 result=None,556 error=ToolError(557 error_type=ToolErrorType.EXECUTION_ERROR,558 message=str(e),559 ),560 )561 562 try:563 result = await asyncio.wait_for(564 self._async_call_tool(action.tool_name, action.arguments),565 timeout=timeout,566 )567 return CallToolObservation(tool_name=action.tool_name, result=result)568 except asyncio.TimeoutError:569 return CallToolObservation(570 tool_name=action.tool_name,571 result=None,572 error=ToolError(573 error_type=ToolErrorType.TIMEOUT,574 message=f"Tool '{action.tool_name}' timed out after {timeout} seconds",575 ),576 )577 except Exception as e:578 error_message = str(e)579 if (580 "not found" in error_message.lower()581 or "unknown tool" in error_message.lower()582 ):583 error_type = ToolErrorType.TOOL_NOT_FOUND584 elif (585 "invalid" in error_message.lower()586 or "argument" in error_message.lower()587 ):588 error_type = ToolErrorType.INVALID_ARGS589 else:590 error_type = ToolErrorType.EXECUTION_ERROR591 return CallToolObservation(592 tool_name=action.tool_name,593 result=None,594 error=ToolError(error_type=error_type, message=error_message),595 )596 597 async def step_async(598 self,599 action: Action,600 timeout_s: Optional[float] = None,601 **kwargs: Any,602 ) -> Observation:603 """604 Async step that routes MCP actions without going through run_async_safely.605 606 The WebSocket handler calls this directly on the outer event loop, where607 the MCP session is already open, avoiding the thread/event-loop deadlock608 that occurs when the sync step() path is used via run_in_executor.609 """610 if isinstance(action, ListToolsAction):611 return await self._async_handle_list_tools()612 elif isinstance(action, CallToolAction):613 return await self._async_handle_call_tool(action, timeout_s=timeout_s)614 else:615 loop = asyncio.get_event_loop()616 return await loop.run_in_executor(617 None, lambda: self._step_impl(action, timeout_s=timeout_s, **kwargs)618 )619 620 @abstractmethod621 def _step_impl(622 self,623 action: Action,624 timeout_s: Optional[float] = None,625 **kwargs: Any,626 ) -> Observation:627 """628 Handle non-MCP actions in the environment.629 630 Subclasses must implement this method to handle any actions that are631 not ListToolsAction or CallToolAction. This is where environment-specific632 action processing should occur.633 634 Args:635 action: The action to execute (guaranteed not to be an MCP action).636 timeout_s: Optional timeout in seconds.637 **kwargs: Additional arguments.638 639 Returns:640 An Observation appropriate for the action.641 """642 pass643 644 def close(self) -> None:645 """646 Clean up resources used by the environment.647 648 This method cleans up the MCP client and any other resources.649 Subclasses should call super().close() if they override this method.650 """651 # The MCP client uses async context manager, so cleanup happens652 # automatically when the context exits. We just clear references.653 self.mcp_client = None654 self.mcp_server = None655 