Team Ai
Apppublic

openenv/coding_env

sourceHugging Faceupdated 3mo agoView on Hugging Face
21likes
mcp_environment.py646 linesDownload Raw Back to env_server
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