openenv/coding_env
21
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 