Team Ai
Apppublic

openenv/echo_env

sourceHugging Faceupdated 1d agoView on Hugging Face
6likes
mcp_environment.py655 linesDownload Raw Back to env_server
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