Team Ai
Apppublic

openenv/chat_env

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
mcp_client.py485 linesDownload Raw Back to core
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 Client classes for tool-calling environments.9 10This module provides async client classes for interacting with MCP-enabled environments:11- MCPClientBase: Base class with shared tool discovery12- MCPToolClient: Client for tool-calling style (one tool per step)13 14These clients abstract away the MCP protocol details, providing a clean interface15for listing and calling tools on remote environments. All clients are async by default.16 17Architecture Overview::18 19    ┌─────────────────────────────────────────────────────────┐20    │                    HTTPEnvServer                        │21    ├─────────────────────────────────────────────────────────┤22    │  Simulation Mode (default):                             │23    │    /ws    → OpenEnv protocol (reset/step/state)         │24    │    /mcp   → MCP JSON-RPC (tools/list, tools/call)       │25    │    /reset, /step, /state → HTTP endpoints               │26    ├─────────────────────────────────────────────────────────┤27    │  Production Mode (use_production_mode=True):                     │28    │    /mcp   → MCP JSON-RPC (tools/list, tools/call)       │29    │    Bypasses step() for direct tool access               │30    └─────────────────────────────────────────────────────────┘31 32    Client Usage:33      MCPToolClient (default)     → /ws (step-based, with rewards)34      MCPToolClient (production)    → /mcp (direct tool access, no rewards)35 36Example (async):37    >>> from openenv.core.mcp_client import MCPToolClient38    >>>39    >>> async with MCPToolClient(base_url="http://localhost:8000") as env:40    ...     # Discover available tools41    ...     tools = await env.list_tools()42    ...     print([t.name for t in tools])43    ...44    ...     # Call a tool45    ...     result = await env.call_tool("echo_message", message="Hello!")46    ...     print(result)47 48Example (sync wrapper):49    >>> env = MCPToolClient(base_url="http://localhost:8000").sync()50    >>> with env:51    ...     tools = env.list_tools()52    ...     result = env.call_tool("echo_message", message="Hello!")53"""54 55import asyncio56from typing import Any, Dict, List, Optional57 58from .client_types import StepResult59from .env_client import EnvClient60from .env_server.mcp_types import (61    CallToolAction,62    CallToolObservation,63    ListToolsAction,64    ListToolsObservation,65    Tool,66    ToolError,67)68from .env_server.types import Observation, State69 70 71class MCPClientBase(EnvClient[Any, Observation, State]):72    """73    Base class for MCP clients with tool discovery.74 75    This class provides the common `list_tools()` method for discovering76    available tools from an MCP-enabled environment. Subclasses implement77    specific interaction patterns (tool-calling or CodeAct).78 79    Attributes:80        _tools_cache: Cached list of tools (populated on first `list_tools()` call)81    """82 83    def __init__(84        self,85        base_url: str,86        connect_timeout_s: float = 10.0,87        message_timeout_s: float = 60.0,88        provider: Optional[Any] = None,89        mode: Optional[str] = None,90    ):91        """92        Initialize MCP client.93 94        Args:95            base_url: Base URL of the environment server (http:// or ws://).96            connect_timeout_s: Timeout for establishing WebSocket connection.97            message_timeout_s: Timeout for receiving responses to messages.98            provider: Optional container/runtime provider for lifecycle management.99            mode: Communication mode. Must be 'production' for MCP clients. Defaults to 'production'.100        """101        # MCPClientBase defaults to production mode, but allow override for validation102        if mode is None:103            mode = "production"104 105        # Validate that mode is production106        mode_lower = mode.lower()107        if mode_lower != "production":108            raise ValueError(109                f"MCPToolClient only supports 'production' mode, got '{mode}'. "110                f"Use GenericEnvClient for simulation mode."111            )112 113        super().__init__(114            base_url=base_url,115            connect_timeout_s=connect_timeout_s,116            message_timeout_s=message_timeout_s,117            provider=provider,118            mode=mode,119        )120        self._tools_cache: Optional[List[Tool]] = None121        self.use_production_mode = False122        self._production_session_id: Optional[str] = None123        self._production_session_lock = asyncio.Lock()124        self._jsonrpc_request_id = 0125        self._http_client: Optional[Any] = None  # lazily-created httpx.AsyncClient126 127    def _next_request_id(self) -> int:128        """Generate a monotonically increasing JSON-RPC request id."""129        self._jsonrpc_request_id += 1130        return self._jsonrpc_request_id131 132    def _production_mcp_url(self) -> str:133        """Build HTTP MCP endpoint URL from the client's websocket URL."""134        url = self._ws_url.replace("ws://", "http://").replace("wss://", "https://")135        if url.endswith("/ws"):136            url = url[: -len("/ws")]137        return url.rstrip("/") + "/mcp"138 139    async def _get_http_client(self) -> Any:140        """Return a shared httpx.AsyncClient, creating one lazily."""141        if self._http_client is None:142            import httpx143 144            self._http_client = httpx.AsyncClient()145        return self._http_client146 147    async def _production_mcp_request(148        self, method: str, params: Optional[Dict[str, Any]] = None149    ) -> Dict[str, Any]:150        """Send a JSON-RPC request to HTTP /mcp and return parsed JSON response."""151        client = await self._get_http_client()152        response = await client.post(153            self._production_mcp_url(),154            json={155                "jsonrpc": "2.0",156                "method": method,157                "params": params or {},158                "id": self._next_request_id(),159            },160            timeout=self._message_timeout,161        )162        response.raise_for_status()163        return response.json()164 165    async def _ensure_production_session(self) -> str:166        """Create and cache a persistent HTTP MCP session id if needed."""167        async with self._production_session_lock:168            if self._production_session_id is not None:169                return self._production_session_id170 171            data = await self._production_mcp_request("openenv/session/create")172            if "error" in data:173                message = data.get("error", {}).get("message", "unknown error")174                raise RuntimeError(f"Failed to create MCP session: {message}")175 176            session_id = data.get("result", {}).get("session_id")177            if not session_id:178                raise RuntimeError("Failed to create MCP session: missing session_id")179 180            self._production_session_id = session_id181            return session_id182 183    async def list_tools(self, use_cache: bool = True) -> List[Tool]:184        """185        Discover available tools from the environment.186 187        Args:188            use_cache: If True, return cached tools if available.189                      Set to False to force a fresh request.190 191        Returns:192            List of Tool objects with name, description, and input_schema.193 194        Example:195            >>> tools = await env.list_tools()196            >>> for tool in tools:197            ...     print(f"{tool.name}: {tool.description}")198        """199        if use_cache and self._tools_cache is not None:200            return self._tools_cache201 202        # Use production mode HTTP endpoint if enabled.203        # Some tests instantiate with __new__ and skip __init__, so default missing flag to False.204        if getattr(self, "use_production_mode", False):205            try:206                session_id = await self._ensure_production_session()207                data = await self._production_mcp_request(208                    "tools/list",209                    {"session_id": session_id},210                )211                if "error" in data:212                    message = data.get("error", {}).get("message", "unknown error")213                    raise RuntimeError(f"list_tools failed: {message}")214                if "result" in data and "tools" in data["result"]:215                    tools = [216                        Tool(217                            name=t.get("name", ""),218                            description=t.get("description", ""),219                            input_schema=t.get(220                                "input_schema", t.get("inputSchema", {})221                            ),222                        )223                        for t in data["result"]["tools"]224                    ]225                    self._tools_cache = tools226                    return tools227            except Exception:228                # If HTTP request fails, return empty list229                pass230            return []231 232        result = await self.step(ListToolsAction())233        if isinstance(result.observation, ListToolsObservation):234            self._tools_cache = result.observation.tools235            return self._tools_cache236 237        # Unexpected observation type; keep API stable with an empty tool list.238        self._tools_cache = []239        return self._tools_cache240 241    def _step_payload(self, action: Any) -> Dict[str, Any]:242        """Convert an Action object to the JSON data expected by the env server."""243        if isinstance(action, ListToolsAction):244            return {"type": "list_tools"}245        elif isinstance(action, CallToolAction):246            return {247                "type": "call_tool",248                "tool_name": action.tool_name,249                "arguments": action.arguments,250            }251        else:252            # For unknown actions, try to serialize as dict253            if hasattr(action, "model_dump"):254                return action.model_dump()255            return {"action": str(action)}256 257    def _parse_result(self, payload: Dict[str, Any]) -> StepResult[Observation]:258        """Convert a JSON response from the env server to StepResult[Observation]."""259        obs_data = payload.get("observation", {})260 261        # Check if this is a ListToolsObservation262        if "tools" in obs_data:263            tools = [264                Tool(265                    name=t.get("name", ""),266                    description=t.get("description", ""),267                    input_schema=t.get("input_schema", t.get("inputSchema", {})),268                )269                for t in obs_data.get("tools", [])270            ]271            observation = ListToolsObservation(272                tools=tools,273                done=payload.get("done", False),274                reward=payload.get("reward"),275                metadata=obs_data.get("metadata", {}),276            )277        # Check if this is a CallToolObservation278        elif "tool_name" in obs_data:279            error = None280            if obs_data.get("error"):281                error = ToolError(**obs_data["error"])282 283            observation = CallToolObservation(284                tool_name=obs_data.get("tool_name", ""),285                result=obs_data.get("result"),286                error=error,287                done=payload.get("done", False),288                reward=payload.get("reward"),289                metadata=obs_data.get("metadata", {}),290            )291        else:292            # Generic observation293            observation = Observation(294                done=payload.get("done", False),295                reward=payload.get("reward"),296                metadata=obs_data.get("metadata", {}),297            )298 299        return StepResult(300            observation=observation,301            reward=payload.get("reward"),302            done=payload.get("done", False),303        )304 305    def _parse_state(self, payload: Dict[str, Any]) -> State:306        """Convert a JSON response from the state endpoint to a State object."""307        return State(308            episode_id=payload.get("episode_id"),309            step_count=payload.get("step_count", 0),310        )311 312    async def close(self) -> None:313        """314        Close client resources.315 316        In production MCP mode, this also closes the server-side persistent317        MCP session (best effort) before closing websocket/provider resources.318        """319        if self._production_session_id is not None:320            try:321                await self._production_mcp_request(322                    "openenv/session/close",323                    {"session_id": self._production_session_id},324                )325            except Exception:326                # Best effort cleanup - do not mask normal close behavior327                pass328            finally:329                self._production_session_id = None330 331        if self._http_client is not None:332            try:333                await self._http_client.aclose()334            except Exception:335                pass336            finally:337                self._http_client = None338 339        await super().close()340 341 342class MCPToolClient(MCPClientBase):343    """344    Async client for tool-calling style MCP interactions.345 346    Each step invokes a single tool. Use this for traditional function-calling347    agent patterns where the agent decides which tool to call next.348 349    This client provides convenience methods for tool discovery and invocation:350    - `list_tools()`: Get all available tools with their schemas351    - `call_tool(name, **kwargs)`: Invoke a tool by name with arguments352 353    Example (async):354        >>> async with MCPToolClient(base_url="http://localhost:8000") as env:355        ...     # Reset the environment356        ...     await env.reset()357        ...358        ...     # Discover available tools359        ...     tools = await env.list_tools()360        ...     print([t.name for t in tools])  # ['echo_message', 'echo_with_length']361        ...362        ...     # Call a tool directly363        ...     result = await env.call_tool("echo_message", message="Hello!")364        ...     print(result)  # "Hello!"365        ...366        ...     # Or use the full action interface367        ...     from openenv.core.env_server.mcp_types import CallToolAction368        ...     step_result = await env.step(CallToolAction(369        ...         tool_name="echo_with_length",370        ...         arguments={"message": "Test"}371        ...     ))372        ...     print(step_result.observation.result)373 374    Example (sync wrapper):375        >>> env = MCPToolClient(base_url="http://localhost:8000").sync()376        >>> with env:377        ...     tools = env.list_tools()378        ...     result = env.call_tool("echo_message", message="Hello!")379    """380 381    async def call_tool(self, name: str, **kwargs: Any) -> Any:382        """383        Call a tool by name.384 385        This is a convenience method that creates a CallToolAction, executes it,386        and returns the result directly. For more control, use `step()` with387        a CallToolAction directly.388 389        Args:390            name: Name of the tool to invoke (must match a tool from `list_tools()`).391            **kwargs: Arguments to pass to the tool. Must match the tool's input_schema.392 393        Returns:394            The tool's result. The type depends on the tool being called.395 396        Raises:397            RuntimeError: If the server returns an error response.398 399        Example:400            >>> result = await env.call_tool("add", a=5, b=3)401            >>> print(result)  # 8402            >>>403            >>> result = await env.call_tool("greet", name="Claude")404            >>> print(result)  # "Hello, Claude!"405        """406        if getattr(self, "use_production_mode", False):407            session_id = await self._ensure_production_session()408            data = await self._production_mcp_request(409                "tools/call",410                {411                    "name": name,412                    "arguments": kwargs,413                    "session_id": session_id,414                },415            )416 417            if "error" in data:418                message = data.get("error", {}).get("message", "unknown error")419                raise RuntimeError(f"Tool '{name}' failed: {message}")420 421            result = data.get("result")422            if isinstance(result, dict) and "data" in result:423                return result["data"]424            return result425 426        action = CallToolAction(tool_name=name, arguments=kwargs)427        result = await self.step(action)428        obs = result.observation429 430        # Check for transport/framework errors431        if isinstance(obs, CallToolObservation) and obs.error is not None:432            raise RuntimeError(433                f"Tool '{name}' failed: {obs.error.message} "434                f"(type: {obs.error.error_type.value})"435            )436 437        # Return the result438        if isinstance(obs, CallToolObservation):439            result = obs.result440            # Handle FastMCP CallToolResult objects441            # - As object: has .data attribute442            # - As dict (from JSON): has "data" key443            if hasattr(result, "data"):444                return result.data445            if isinstance(result, dict) and "data" in result:446                return result["data"]447            return result448 449        # Fallback for unexpected observation types450        return obs451 452    async def get_tool(self, name: str) -> Optional[Tool]:453        """454        Get a specific tool by name.455 456        Args:457            name: Name of the tool to find.458 459        Returns:460            The Tool object if found, None otherwise.461 462        Example:463            >>> tool = await env.get_tool("echo_message")464            >>> if tool:465            ...     print(tool.description)466            ...     print(tool.input_schema)467        """468        tools = await self.list_tools()469        for tool in tools:470            if tool.name == name:471                return tool472        return None473 474    async def has_tool(self, name: str) -> bool:475        """476        Check if a tool exists.477 478        Args:479            name: Name of the tool to check.480 481        Returns:482            True if the tool exists, False otherwise.483        """484        return await self.get_tool(name) is not None485