openenv/chat_env
0
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 Environment base class for OpenEnv.9 10This module provides the MCPEnvironment base class that integrates FastMCP servers11with OpenEnv's Gym-style Environment interface. It handles MCP tool discovery12and invocation through the step() API, following RFC 003.13 14Key features:15- Automatic routing of ListToolsAction and CallToolAction to MCP server16- Reserved tool name validation (reset, step, state, close are protected)17- Timeout handling for tool calls18- Proper error categorization (tool not found, execution errors, timeouts)19- Mode-aware tool registration (production vs simulation)20- Code mode support via get_callables() and execute_code()21 22Usage:23 from fastmcp import FastMCP24 from openenv.core.env_server.mcp_environment import MCPEnvironment25 26 class MyMCPEnv(MCPEnvironment):27 def __init__(self):28 mcp = FastMCP("my-server")29 30 # Register mode-specific tools31 @self.tool(mode="production")32 def my_tool(arg: str) -> str:33 return f"Production: {arg}"34 35 @self.tool(mode="simulation")36 def my_tool(arg: str) -> str:37 return f"Simulation: {arg}"38 39 super().__init__(mcp)40 41 def reset(self, seed=None, episode_id=None, **kwargs):42 # Reset logic here43 ...44 45 def _step_impl(self, action):46 # Handle non-MCP actions47 ...48 49 @property50 def state(self):51 # Return current state52 ...53"""54 55import asyncio56import inspect57from abc import abstractmethod58from collections import defaultdict59from contextlib import asynccontextmanager60from typing import Any, Callable, Dict, Optional61 62from fastmcp import Client63from fastmcp.client.client import CallToolResult64from mcp.types import TextContent65 66from ..utils import run_async_safely67from .interfaces import Environment68from .mcp_types import (69 CallToolAction,70 CallToolObservation,71 ListToolsAction,72 ListToolsObservation,73 RESERVED_TOOL_NAMES,74 Tool,75 ToolError,76 ToolErrorType,77)78from .types import Action, Observation79 80 81# Default timeout for MCP tool calls in seconds82MCP_TOOL_CALL_TIMEOUT = 30.083 84# Valid modes for tool registration85VALID_MODES = {"production", "simulation"}86 87 88def get_server_tools(mcp_server: Any) -> Dict[str, Any]:89 """90 Get tools from a FastMCP server, compatible with both 2.x and 3.x.91 92 Returns:93 Dictionary mapping tool names to tool objects.94 """95 # FastMCP 2.x: get_tools() returns dict {name: Tool}96 if hasattr(mcp_server, "get_tools"):97 result = run_async_safely(mcp_server.get_tools())98 if isinstance(result, dict):99 return result100 # FastMCP 3.x: list_tools() returns list of Tool objects101 if hasattr(mcp_server, "list_tools"):102 tools_list = run_async_safely(mcp_server.list_tools())103 return {t.name: t for t in tools_list}104 return {}105 106 107class MCPEnvironment(Environment):108 """109 Base class for environments that expose tools via MCP (Model Context Protocol).110 111 MCPEnvironment bridges FastMCP servers with OpenEnv's Gym-style API, allowing112 agents to discover and invoke MCP tools through the standard step() interface.113 114 The class automatically handles:115 - ListToolsAction: Returns available tools from the MCP server116 - CallToolAction: Invokes a specific tool with arguments117 118 All other actions are delegated to the abstract _step_impl() method,119 which subclasses must implement.120 121 Args:122 mcp_server: A FastMCP server instance containing tool definitions.123 The server's tools will be validated against reserved names.124 transform: Optional transform to apply to observations (inherited from Environment).125 126 Raises:127 ValueError: If any tool in the MCP server uses a reserved name128 (reset, step, state, close).129 130 Example:131 >>> from fastmcp import FastMCP132 >>> mcp = FastMCP("calculator")133 >>> @mcp.tool()134 ... def add(a: int, b: int) -> int:135 ... return a + b136 >>> env = MyMCPEnvironment(mcp)137 >>> obs = env.step(ListToolsAction())138 >>> obs.tools[0].name139 'add'140 """141 142 def __init__(self, mcp_server: Any, transform: Optional[Any] = None) -> None:143 """144 Initialize the MCP environment.145 146 Args:147 mcp_server: A FastMCP server instance with tool definitions.148 transform: Optional transform to apply to observations.149 150 Raises:151 ValueError: If any tool uses a reserved name (reset, step, state, close).152 """153 super().__init__(transform=transform)154 155 # Validate tool names before storing156 self._validate_tool_names(mcp_server)157 158 self.mcp_server = mcp_server159 self.mcp_client = Client(mcp_server)160 161 # Track mode-specific tools: {tool_name: {mode: func}}162 # mode can be "production", "simulation", or None (available in all modes)163 self._mode_tools = defaultdict(dict)164 165 # Track tool schemas for list_tools: {tool_name: {mode: schema}}166 self._mode_tool_schemas = defaultdict(dict)167 168 def _require_mcp_client(self) -> Any:169 """Return MCP client or raise if environment has been closed."""170 if self.mcp_client is None:171 raise RuntimeError("MCP client is not available; environment is closed")172 return self.mcp_client173 174 def _require_mcp_server(self) -> Any:175 """Return MCP server or raise if environment has been closed."""176 if self.mcp_server is None:177 raise RuntimeError("MCP server is not available; environment is closed")178 return self.mcp_server179 180 @asynccontextmanager181 async def mcp_session(self):182 """183 Context manager for MCP client sessions.184 185 This wrapper serves two purposes:186 187 1. **Null guard** — raises a clear error if ``close()`` has already188 been called (``mcp_client`` is ``None``).189 190 2. **AsyncExitStack adapter** — FastMCP's ``Client.__aenter__``191 creates a background ``asyncio.Task`` for session management.192 When entered directly via ``AsyncExitStack`` in the HTTP session193 path (``_create_session``), this task can be cancelled by ASGI194 harnesses (e.g. Starlette ``TestClient``) between requests,195 corrupting session state. Wrapping in an ``asynccontextmanager``196 generator isolates the task lifecycle: the generator frame keeps197 ``async with client:`` suspended at ``yield``, so cleanup only198 runs when the stack explicitly closes the generator — not when199 the event loop cancels orphaned tasks.200 201 Delegates to FastMCP's ``Client`` context manager which is202 reentrant: the first entry opens the transport and subsequent203 (nested) entries simply increment an internal reference counter.204 The transport is closed only when the outermost context exits.205 206 No external lock is needed because ``Client._connect`` /207 ``Client._disconnect`` already serialise connection state changes208 through their own ``anyio.Lock``.209 """210 client = self._require_mcp_client()211 async with client:212 yield client213 214 @property215 def supports_code_mode(self) -> bool:216 """Check if this environment supports code mode (execute_code)."""217 return True218 219 def _get_server_tools(self, mcp_server: Any) -> Dict[str, Any]:220 """221 Get tools from a FastMCP server, compatible with both 2.x and 3.x.222 223 Returns:224 Dictionary mapping tool names to tool objects.225 """226 return get_server_tools(mcp_server)227 228 def get_callables(self) -> Dict[str, Callable]:229 """230 Get callable functions for code mode.231 232 Returns tool functions as direct Python callables, enabling code mode233 where agents write Python code that calls tools directly (no JSON-RPC234 overhead). Mode-specific tools are filtered by the current mode.235 236 Returns:237 Dictionary mapping tool names to callables.238 """239 callables: Dict[str, Callable] = {}240 current_mode = getattr(self, "_mode", None)241 242 # Extract callables from FastMCP server using public API243 for tool_name, tool in self._get_server_tools(self.mcp_server).items():244 if hasattr(tool, "fn") and callable(tool.fn):245 callables[tool_name] = tool.fn246 247 # Add mode-specific tools available in current mode248 for tool_name, mode_funcs in self._mode_tools.items():249 if None in mode_funcs:250 # Tool available in all modes (already in FastMCP if registered there)251 if tool_name not in callables:252 callables[tool_name] = mode_funcs[None]253 elif current_mode in mode_funcs:254 # Tool available in current mode only255 callables[tool_name] = mode_funcs[current_mode]256 257 return callables258 259 def execute_code(self, code: str) -> Observation:260 """261 Execute Python code with tools available as callables.262 263 This enables the CodeAct pattern where agents write Python code264 that calls tools directly as functions, avoiding JSON-RPC overhead.265 266 Args:267 code: Python code to execute. Tools are available as functions268 in the execution namespace. Set a variable named 'result'269 to capture the return value.270 271 Returns:272 Observation with result in metadata["result"] or error in273 metadata["error"].274 """275 namespace = self.get_callables()276 277 result_dict: Dict[str, Any] = {}278 try:279 exec(code, namespace, result_dict)280 result = result_dict.get("result")281 return Observation(done=False, reward=0.0, metadata={"result": result})282 except SyntaxError as e:283 return Observation(284 done=False, reward=0.0, metadata={"error": f"Syntax error: {str(e)}"}285 )286 except Exception as e:287 return Observation(done=False, reward=0.0, metadata={"error": str(e)})288 289 def _validate_tool_names(self, mcp_server: Any) -> None:290 """291 Validate that no tools use reserved names.292 293 Reserved names (reset, step, state, close) are protected to maintain294 the dual API boundary between infrastructure and agent APIs.295 296 Args:297 mcp_server: The FastMCP server to validate.298 299 Raises:300 ValueError: If any tool uses a reserved name.301 """302 tools_dict = self._get_server_tools(mcp_server)303 if tools_dict:304 tool_names = set(tools_dict.keys())305 conflicts = tool_names & RESERVED_TOOL_NAMES306 if conflicts:307 raise ValueError(308 f"MCP tools cannot use reserved names: {sorted(conflicts)}. "309 f"Reserved names are: {sorted(RESERVED_TOOL_NAMES)}"310 )311 312 def tool(self, mode: Optional[str] = None) -> Callable:313 """314 Decorator for registering mode-aware tools.315 316 Args:317 mode: Optional mode for the tool ("production" or "simulation").318 If None, tool is available in all modes.319 320 Returns:321 A decorator function for registering tools.322 323 Raises:324 ValueError: If mode is not None, "production", or "simulation".325 """326 if mode is not None and mode not in VALID_MODES:327 raise ValueError(328 f"Invalid mode '{mode}'. Mode must be 'production', 'simulation', or None."329 )330 331 def decorator(func: Callable) -> Callable:332 tool_name = func.__name__333 # Validate tool name is not reserved334 if tool_name in RESERVED_TOOL_NAMES:335 raise ValueError(336 f"Tool name '{tool_name}' is reserved and cannot be used. "337 f"Reserved names are: {sorted(RESERVED_TOOL_NAMES)}"338 )339 340 # If mode is None, register with FastMCP as usual341 if mode is None:342 mcp_server = self._require_mcp_server()343 decorated_func = mcp_server.tool()(func)344 self._mode_tools[tool_name][None] = func345 return decorated_func346 347 # For mode-specific tools, don't register with FastMCP348 # Instead, track them ourselves349 self._mode_tools[tool_name][mode] = func350 351 # Extract schema information from function signature352 sig = inspect.signature(func)353 schema = {354 "type": "object",355 "properties": {},356 "required": [],357 }358 359 for param_name, param in sig.parameters.items():360 # Get type annotation361 param_type = param.annotation362 json_type = "string" # default363 if param_type in (int, "int"):364 json_type = "integer"365 elif param_type in (float, "float"):366 json_type = "number"367 elif param_type in (bool, "bool"):368 json_type = "boolean"369 370 schema["properties"][param_name] = {"type": json_type}371 372 # If no default value, it's required373 if param.default == inspect.Parameter.empty:374 schema["required"].append(param_name)375 376 # Store the schema for this mode-specific tool377 self._mode_tool_schemas[tool_name][mode] = {378 "name": tool_name,379 "description": func.__doc__ or "",380 "input_schema": schema,381 }382 383 return func384 385 return decorator386 387 def step(388 self,389 action: Action,390 timeout_s: Optional[float] = None,391 **kwargs: Any,392 ) -> Observation:393 """394 Execute an action in the environment.395 396 This method routes MCP-specific actions (ListToolsAction, CallToolAction)397 to the appropriate handlers, while delegating all other actions to398 the subclass's _step_impl() method.399 400 Args:401 action: The action to execute. Can be:402 - ListToolsAction: Returns available MCP tools403 - CallToolAction: Invokes a specific MCP tool404 - Any other Action: Delegated to _step_impl()405 timeout_s: Optional timeout in seconds for the action.406 Defaults to MCP_TOOL_CALL_TIMEOUT (30s) for MCP actions.407 **kwargs: Additional arguments passed to handlers.408 409 Returns:410 Observation appropriate to the action type:411 - ListToolsObservation for ListToolsAction412 - CallToolObservation for CallToolAction413 - Subclass-defined Observation for other actions414 """415 if isinstance(action, ListToolsAction):416 return self._handle_list_tools()417 elif isinstance(action, CallToolAction):418 return self._handle_call_tool(action, timeout_s=timeout_s)419 else:420 return self._step_impl(action, timeout_s=timeout_s, **kwargs)421 422 def _handle_list_tools(self) -> ListToolsObservation:423 """Sync wrapper — delegates to the canonical async implementation."""424 return run_async_safely(self._async_handle_list_tools())425 426 async def _async_list_tools(self) -> list:427 """428 Async helper to list tools from the MCP client.429 430 Returns:431 List of tool objects from the MCP server.432 """433 async with self.mcp_session() as client:434 return await client.list_tools()435 436 def _handle_call_tool(437 self,438 action: CallToolAction,439 timeout_s: Optional[float] = None,440 ) -> CallToolObservation:441 """Sync wrapper — delegates to the canonical async implementation."""442 return run_async_safely(443 self._async_handle_call_tool(action, timeout_s=timeout_s)444 )445 446 async def _async_call_tool(self, tool_name: str, arguments: dict) -> Any:447 """448 Async helper to call a tool on the MCP server.449 450 Args:451 tool_name: Name of the tool to invoke.452 arguments: Dictionary of arguments to pass to the tool.453 454 Returns:455 The result from the tool execution.456 """457 async with self.mcp_session() as client:458 return await client.call_tool(tool_name, arguments)459 460 async def _async_handle_list_tools(self) -> ListToolsObservation:461 """Async version of _handle_list_tools — avoids run_async_safely."""462 try:463 current_mode = getattr(self, "_mode", None)464 tools_result = await self._async_list_tools()465 tools = []466 for tool in tools_result:467 if tool.name not in self._mode_tool_schemas:468 tools.append(469 Tool(470 name=tool.name,471 description=tool.description or "",472 input_schema=tool.inputSchema473 if hasattr(tool, "inputSchema")474 else {},475 )476 )477 for tool_name, mode_schemas in self._mode_tool_schemas.items():478 if None in mode_schemas:479 schema = mode_schemas[None]480 tools.append(481 Tool(482 name=schema["name"],483 description=schema["description"],484 input_schema=schema["input_schema"],485 )486 )487 elif current_mode in mode_schemas:488 schema = mode_schemas[current_mode]489 tools.append(490 Tool(491 name=schema["name"],492 description=schema["description"],493 input_schema=schema["input_schema"],494 )495 )496 return ListToolsObservation(tools=tools)497 except Exception as e:498 return ListToolsObservation(499 tools=[],500 metadata={"error": str(e), "error_type": "list_tools_failed"},501 )502 503 async def _async_handle_call_tool(504 self,505 action: CallToolAction,506 timeout_s: Optional[float] = None,507 ) -> CallToolObservation:508 """Async version of _handle_call_tool — avoids run_async_safely."""509 timeout = timeout_s if timeout_s is not None else MCP_TOOL_CALL_TIMEOUT510 tool_name = action.tool_name511 current_mode = getattr(self, "_mode", None)512 513 if tool_name in self._mode_tools:514 mode_info = self._mode_tools[tool_name]515 if None in mode_info:516 func = mode_info[None]517 elif current_mode in mode_info:518 func = mode_info[current_mode]519 else:520 return CallToolObservation(521 tool_name=tool_name,522 result=None,523 error=ToolError(524 error_type=ToolErrorType.TOOL_NOT_FOUND,525 message=f"Tool '{tool_name}' not available in {current_mode} mode",526 ),527 )528 try:529 if inspect.iscoroutinefunction(func):530 result = await func(**action.arguments)531 else:532 result = func(**action.arguments)533 return CallToolObservation(534 tool_name=tool_name,535 result=CallToolResult(536 content=[TextContent(type="text", text=str(result))],537 structured_content={"result": result},538 meta=None,539 data=result,540 is_error=False,541 ),542 )543 except Exception as e:544 return CallToolObservation(545 tool_name=tool_name,546 result=None,547 error=ToolError(548 error_type=ToolErrorType.EXECUTION_ERROR,549 message=str(e),550 ),551 )552 553 try:554 result = await asyncio.wait_for(555 self._async_call_tool(action.tool_name, action.arguments),556 timeout=timeout,557 )558 return CallToolObservation(tool_name=action.tool_name, result=result)559 except asyncio.TimeoutError:560 return CallToolObservation(561 tool_name=action.tool_name,562 result=None,563 error=ToolError(564 error_type=ToolErrorType.TIMEOUT,565 message=f"Tool '{action.tool_name}' timed out after {timeout} seconds",566 ),567 )568 except Exception as e:569 error_message = str(e)570 if (571 "not found" in error_message.lower()572 or "unknown tool" in error_message.lower()573 ):574 error_type = ToolErrorType.TOOL_NOT_FOUND575 elif (576 "invalid" in error_message.lower()577 or "argument" in error_message.lower()578 ):579 error_type = ToolErrorType.INVALID_ARGS580 else:581 error_type = ToolErrorType.EXECUTION_ERROR582 return CallToolObservation(583 tool_name=action.tool_name,584 result=None,585 error=ToolError(error_type=error_type, message=error_message),586 )587 588 async def step_async(589 self,590 action: Action,591 timeout_s: Optional[float] = None,592 **kwargs: Any,593 ) -> Observation:594 """595 Async step that routes MCP actions without going through run_async_safely.596 597 The WebSocket handler calls this directly on the outer event loop, where598 the MCP session is already open, avoiding the thread/event-loop deadlock599 that occurs when the sync step() path is used via run_in_executor.600 """601 if isinstance(action, ListToolsAction):602 return await self._async_handle_list_tools()603 elif isinstance(action, CallToolAction):604 return await self._async_handle_call_tool(action, timeout_s=timeout_s)605 else:606 loop = asyncio.get_event_loop()607 return await loop.run_in_executor(608 None, lambda: self._step_impl(action, timeout_s=timeout_s, **kwargs)609 )610 611 @abstractmethod612 def _step_impl(613 self,614 action: Action,615 timeout_s: Optional[float] = None,616 **kwargs: Any,617 ) -> Observation:618 """619 Handle non-MCP actions in the environment.620 621 Subclasses must implement this method to handle any actions that are622 not ListToolsAction or CallToolAction. This is where environment-specific623 action processing should occur.624 625 Args:626 action: The action to execute (guaranteed not to be an MCP action).627 timeout_s: Optional timeout in seconds.628 **kwargs: Additional arguments.629 630 Returns:631 An Observation appropriate for the action.632 """633 pass634 635 def close(self) -> None:636 """637 Clean up resources used by the environment.638 639 This method cleans up the MCP client and any other resources.640 Subclasses should call super().close() if they override this method.641 """642 # The MCP client uses async context manager, so cleanup happens643 # automatically when the context exits. We just clear references.644 self.mcp_client = None645 self.mcp_server = None646 