Team Ai
Apppublic

openenv/coding_env

sourceHugging Faceupdated 3mo agoView on Hugging Face
21likes
env_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"""8Environment client for persistent sessions.9 10This module provides a WebSocket-based client that maintains a persistent connection11to an environment server, enabling efficient multi-step interactions without12the overhead of HTTP request/response cycles.13 14The client is async by default. For synchronous usage, use the `.sync()` method15to get a `SyncEnvClient` wrapper.16 17Example (async):18    >>> async with GenericEnvClient(base_url="ws://localhost:8000") as env:19    ...     result = await env.reset()20    ...     result = await env.step({"code": "print('hello')"})21 22Example (sync wrapper):23    >>> env = GenericEnvClient(base_url="ws://localhost:8000").sync()24    >>> with env:25    ...     result = env.reset()26    ...     result = env.step({"code": "print('hello')"})27"""28 29from __future__ import annotations30 31import asyncio32import json33import os34from abc import ABC, abstractmethod35from typing import Any, Dict, Generic, Optional, Type, TYPE_CHECKING, TypeVar36 37from .client_types import StateT, StepResult38from .containers.runtime import LocalDockerProvider, UVProvider39from .utils import convert_to_ws_url40 41if TYPE_CHECKING:42    from websockets.asyncio.client import ClientConnection43 44    from .containers.runtime import ContainerProvider, RuntimeProvider45    from .sync_client import SyncEnvClient46 47from websockets.asyncio.client import connect as ws_connect48 49ActT = TypeVar("ActT")50ObsT = TypeVar("ObsT")51EnvClientT = TypeVar("EnvClientT", bound="EnvClient")52 53 54class EnvClient(ABC, Generic[ActT, ObsT, StateT]):55    """56    Async environment client for persistent sessions.57 58    This client maintains a persistent WebSocket connection to an environment59    server, enabling efficient multi-step interactions. Each client instance60    corresponds to a dedicated environment session on the server.61 62    The client is async by default. For synchronous usage, use the `.sync()`63    method to get a `SyncEnvClient` wrapper.64 65    Features:66    - Lower latency for sequential interactions67    - Session state is maintained server-side68    - Better suited for long-running episodes69    - Async by default for modern Python async/await patterns70 71    Example (async):72        >>> from envs.coding_env.client import CodingEnv73        >>>74        >>> # Connect to a server using async context manager75        >>> async with CodingEnv(base_url="ws://localhost:8000") as env:76        ...     result = await env.reset(seed=42)77        ...     while not result.done:78        ...         action = agent.predict(result.observation)79        ...         result = await env.step(action)80 81    Example (sync wrapper):82        >>> env = CodingEnv(base_url="ws://localhost:8000").sync()83        >>> with env:84        ...     result = env.reset(seed=42)85        ...     result = env.step(action)86    """87 88    def __init__(89        self,90        base_url: str,91        connect_timeout_s: float = 10.0,92        message_timeout_s: float = 60.0,93        max_message_size_mb: float = 100.0,94        provider: Optional["ContainerProvider | RuntimeProvider"] = None,95        mode: Optional[str] = None,96    ):97        """98        Initialize environment client.99 100        Args:101            base_url: Base URL of the environment server (http:// or ws://).102                     Will be converted to ws:// if http:// is provided.103            connect_timeout_s: Timeout for establishing WebSocket connection104            message_timeout_s: Timeout for receiving responses to messages105            max_message_size_mb: Maximum WebSocket message size in megabytes.106                                Default 100MB to handle large observations (screenshots, DOM, etc.)107            provider: Optional container/runtime provider for lifecycle management.108                     Can be a ContainerProvider (Docker) or RuntimeProvider (UV).109            mode: Communication mode: 'simulation' for Gym-style API (default) or110                 'production' for MCP JSON-RPC protocol. Can also be set via the111                 OPENENV_CLIENT_MODE environment variable. Constructor parameter112                 takes precedence over environment variable. Case-insensitive.113        """114        # Determine mode (constructor > env var > default)115        if mode is None:116            mode = os.environ.get("OPENENV_CLIENT_MODE", "simulation")117 118        # Normalize and validate mode119        mode = mode.lower()120        if mode not in ("simulation", "production"):121            raise ValueError(122                f"Invalid mode: '{mode}'. Must be 'simulation' or 'production'. "123                f"Set via constructor parameter or OPENENV_CLIENT_MODE environment variable."124            )125 126        # Store mode (use object.__setattr__ to bypass immutability)127        object.__setattr__(self, "_mode", mode)128 129        # Convert HTTP URL to WebSocket URL130        ws_url = convert_to_ws_url(base_url)131 132        self._ws_url = f"{ws_url}/ws"133        self._connect_timeout = connect_timeout_s134        self._message_timeout = message_timeout_s135        self._max_message_size = int(136            max_message_size_mb * 1024 * 1024137        )  # Convert MB to bytes138        self._provider = provider139        self._ws: Optional[ClientConnection] = None140 141    def __setattr__(self, name: str, value: Any) -> None:142        """Prevent modification of _mode after initialization."""143        if name == "_mode" and hasattr(self, "_mode"):144            raise AttributeError("Cannot modify mode after initialization")145        super().__setattr__(name, value)146 147    async def connect(self) -> "EnvClient":148        """149        Establish WebSocket connection to the server.150 151        Returns:152            self for method chaining153 154        Raises:155            ConnectionError: If connection cannot be established156        """157        if self._ws is not None:158            return self159 160        # Bypass proxy for localhost connections161        ws_url_lower = self._ws_url.lower()162        is_localhost = "localhost" in ws_url_lower or "127.0.0.1" in ws_url_lower163 164        old_no_proxy = os.environ.get("NO_PROXY")165        if is_localhost:166            # Set NO_PROXY to bypass proxy for localhost167            current_no_proxy = old_no_proxy or ""168            if "localhost" not in current_no_proxy.lower():169                os.environ["NO_PROXY"] = (170                    f"{current_no_proxy},localhost,127.0.0.1"171                    if current_no_proxy172                    else "localhost,127.0.0.1"173                )174 175        try:176            self._ws = await ws_connect(177                self._ws_url,178                open_timeout=self._connect_timeout,179                max_size=self._max_message_size,180            )181        except Exception as e:182            raise ConnectionError(f"Failed to connect to {self._ws_url}: {e}") from e183        finally:184            # Restore original NO_PROXY value185            if is_localhost:186                if old_no_proxy is None:187                    os.environ.pop("NO_PROXY", None)188                else:189                    os.environ["NO_PROXY"] = old_no_proxy190 191        return self192 193    async def disconnect(self) -> None:194        """Close the WebSocket connection."""195        if self._ws is not None:196            try:197                # Send close message198                await self._send({"type": "close"})199            except Exception:200                pass  # Best effort201            try:202                await self._ws.close()203            except Exception:204                pass205            self._ws = None206 207    async def _ensure_connected(self) -> None:208        """Ensure WebSocket connection is established."""209        if self._ws is None:210            await self.connect()211 212    async def _send(self, message: Dict[str, Any]) -> None:213        """Send a message over the WebSocket."""214        await self._ensure_connected()215        assert self._ws is not None216        await self._ws.send(json.dumps(message))217 218    async def _receive(self) -> Dict[str, Any]:219        """Receive and parse a message from the WebSocket."""220        assert self._ws is not None221        raw = await asyncio.wait_for(self._ws.recv(), timeout=self._message_timeout)222        return json.loads(raw)223 224    async def _send_and_receive(self, message: Dict[str, Any]) -> Dict[str, Any]:225        """Send a message and wait for response."""226        await self._send(message)227        response = await self._receive()228 229        # Check for error response230        if response.get("type") == "error":231            error_data = response.get("data", {})232            raise RuntimeError(233                f"Server error: {error_data.get('message', 'Unknown error')} "234                f"(code: {error_data.get('code', 'UNKNOWN')})"235            )236 237        return response238 239    @classmethod240    async def from_docker_image(241        cls: Type[EnvClientT],242        image: str,243        provider: Optional["ContainerProvider"] = None,244        **kwargs: Any,245    ) -> EnvClientT:246        """247        Create an environment client by spinning up a Docker container.248 249        Args:250            image: Docker image name to run (e.g., "coding-env:latest")251            provider: Container provider to use (defaults to LocalDockerProvider)252            **kwargs: Additional arguments to pass to provider.start_container()253 254        Returns:255            Connected client instance256        """257        if provider is None:258            provider = LocalDockerProvider()259 260        # Start container261        base_url = provider.start_container(image, **kwargs)262 263        # Wait for server to be ready264        provider.wait_for_ready(base_url)265 266        # Create and connect client267        client = cls(base_url=base_url, provider=provider)268        await client.connect()269 270        return client271 272    @classmethod273    async def from_env(274        cls: Type[EnvClientT],275        repo_id: str,276        *,277        use_docker: bool = True,278        provider: Optional["ContainerProvider | RuntimeProvider"] = None,279        **provider_kwargs: Any,280    ) -> EnvClientT:281        """282        Create a client from a Hugging Face Space.283 284        Args:285            repo_id: Hugging Face space identifier ``{org}/{space}``.286            use_docker: When ``True`` (default) pull from the HF registry and287                launch via :class:`LocalDockerProvider`. When ``False`` run the288                space locally with :class:`UVProvider`.289            provider: Optional provider instance to reuse. Must be a290                :class:`ContainerProvider` when ``use_docker=True`` and a291                :class:`RuntimeProvider` otherwise.292            provider_kwargs: Additional keyword arguments forwarded to293                either the container provider's ``start_container`` (docker)294                or to the ``UVProvider`` constructor/start (uv). When295                ``use_docker=False``, the ``project_path`` argument can be296                used to override the default git URL297                (``git+https://huggingface.co/spaces/{repo_id}``).298 299        Returns:300            Connected client instance301 302        Examples:303            >>> # Pull and run from HF Docker registry304            >>> env = await MyEnv.from_env("openenv/echo-env")305            >>>306            >>> # Run locally with UV (clones the space)307            >>> env = await MyEnv.from_env("openenv/echo-env", use_docker=False)308            >>>309            >>> # Run from a local checkout310            >>> env = await MyEnv.from_env(311            ...     "openenv/echo-env",312            ...     use_docker=False,313            ...     project_path="/path/to/local/checkout"314            ... )315        """316        # Extract start args that apply to both providers317        start_args = {}318        for key in ("port", "env_vars", "workers"):319            if key in provider_kwargs:320                start_args[key] = provider_kwargs.pop(key)321 322        if use_docker:323            # Docker mode: pull from HF registry324            docker_provider = provider or LocalDockerProvider()325            tag = provider_kwargs.pop("tag", "latest")326            image = f"registry.hf.space/{repo_id.replace('/', '-')}:{tag}"327            base_url = docker_provider.start_container(328                image, **start_args, **provider_kwargs329            )330            docker_provider.wait_for_ready(base_url)331 332            client = cls(base_url=base_url, provider=docker_provider)333            await client.connect()334            return client335        else:336            # UV mode: clone and run with uv337            if provider is None:338                uv_kwargs = dict(provider_kwargs)339                project_path = uv_kwargs.pop("project_path", None)340                if project_path is None:341                    project_path = f"git+https://huggingface.co/spaces/{repo_id}"342 343                provider = UVProvider(project_path=project_path, **uv_kwargs)344            else:345                if provider_kwargs:346                    raise ValueError(347                        "provider_kwargs cannot be used when supplying a provider instance"348                    )349 350            base_url = provider.start(**start_args)351            provider.wait_for_ready()352 353            client = cls(base_url=base_url, provider=provider)354            await client.connect()355            return client356 357    @abstractmethod358    def _step_payload(self, action: ActT) -> Dict[str, Any]:359        """Convert an Action object to the JSON data expected by the env server."""360        raise NotImplementedError361 362    @abstractmethod363    def _parse_result(self, payload: Dict[str, Any]) -> StepResult[ObsT]:364        """Convert a JSON response from the env server to StepResult[ObsT]."""365        raise NotImplementedError366 367    @abstractmethod368    def _parse_state(self, payload: Dict[str, Any]) -> StateT:369        """Convert a JSON response from the state endpoint to a State object."""370        raise NotImplementedError371 372    async def reset(self, **kwargs: Any) -> StepResult[ObsT]:373        """374        Reset the environment with optional parameters.375 376        Args:377            **kwargs: Optional parameters passed to the environment's reset method.378                     Common parameters include:379                     - seed: Random seed for reproducibility380                     - episode_id: Custom episode identifier381 382        Returns:383            StepResult containing initial observation384        """385        message = {386            "type": "reset",387            "data": kwargs,388        }389        response = await self._send_and_receive(message)390        return self._parse_result(response.get("data", {}))391 392    async def step(self, action: ActT, **kwargs: Any) -> StepResult[ObsT]:393        """394        Execute an action in the environment.395 396        Args:397            action: The action to execute398            **kwargs: Optional parameters (currently ignored)399 400        Returns:401            StepResult containing observation, reward, and done status402        """403        message = {404            "type": "step",405            "data": self._step_payload(action),406        }407        response = await self._send_and_receive(message)408        return self._parse_result(response.get("data", {}))409 410    async def state(self) -> StateT:411        """412        Get the current environment state from the server.413 414        Returns:415            State object with environment state information416        """417        message = {"type": "state"}418        response = await self._send_and_receive(message)419        return self._parse_state(response.get("data", {}))420 421    async def close(self) -> None:422        """423        Close the WebSocket connection and clean up resources.424 425        If this client was created via from_docker_image() or from_env(),426        this will also stop and remove the associated container/process.427        """428        await self.disconnect()429 430        if self._provider is not None:431            # Handle both ContainerProvider and RuntimeProvider432            if hasattr(self._provider, "stop_container"):433                self._provider.stop_container()434            elif hasattr(self._provider, "stop"):435                self._provider.stop()436 437    async def __aenter__(self) -> "EnvClient":438        """Enter async context manager, ensuring connection is established."""439        await self.connect()440        return self441 442    async def __aexit__(self, exc_type, exc_val, exc_tb) -> None:443        """Exit async context manager, closing connection."""444        await self.close()445 446    def __enter__(self) -> "EnvClient":447        """Sync context manager entry - raises error suggesting async usage."""448        raise TypeError(449            "EnvClient is async by default. Use 'async with' instead of 'with', "450            "or call .sync() to get a synchronous wrapper:\n"451            "  async with client:  # async usage\n"452            "  with client.sync():  # sync wrapper"453        )454 455    def __exit__(self, exc_type, exc_val, exc_tb) -> None:456        """Sync context manager exit - should not be reached."""457        pass  # pragma: no cover458 459    def sync(self) -> "SyncEnvClient":460        """461        Return a synchronous wrapper around this async client.462 463        Use this method when you need synchronous access to the environment464        without async/await syntax. This is useful for:465        - Integration with synchronous codebases466        - Interactive/REPL usage467        - Stopping async from "infecting" the call stack468 469        Returns:470            SyncEnvClient wrapper that provides synchronous methods471 472        Example:473            >>> # Create async client and get sync wrapper474            >>> async_client = GenericEnvClient(base_url="http://localhost:8000")475            >>> sync_client = async_client.sync()476            >>>477            >>> # Use synchronous API478            >>> with sync_client:479            ...     result = sync_client.reset()480            ...     result = sync_client.step({"code": "print('hello')"})481        """482        from .sync_client import SyncEnvClient483 484        return SyncEnvClient(self)485