Team Ai
Apppublic

openenv/echo_env

sourceHugging Faceupdated 1d agoView on Hugging Face
6likes
mcp_client.py597 linesDownload Raw Back to core
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