openenv/echo_env
6
1# SPDX-License-Identifier: BSD-3-Clause2 3"""4Environment client for persistent sessions.5 6This module provides a WebSocket-based client that maintains a persistent connection7to an environment server, enabling efficient multi-step interactions without8the overhead of HTTP request/response cycles.9 10The client is async by default. For synchronous usage, use the `.sync()` method11to get a `SyncEnvClient` wrapper.12 13Examples:14 15 Async usage:16 17 ```python18 async with GenericEnvClient(base_url="ws://localhost:8000") as env:19 result = await env.reset()20 result = await env.step({"code": "print('hello')"})21 ```22 23 Sync usage via `.sync()` wrapper:24 25 ```python26 env = GenericEnvClient(base_url="ws://localhost:8000").sync()27 with env:28 result = env.reset()29 result = env.step({"code": "print('hello')"})30 ```31"""32 33from __future__ import annotations34 35import asyncio36import inspect37import ipaddress38import json39import os40import time41from abc import ABC, abstractmethod42from collections.abc import Coroutine43from contextlib import suppress44from typing import Any, Callable, Dict, Generic, Optional, Type, TYPE_CHECKING, TypeVar45from urllib.parse import urlsplit46 47from .client_types import StateT, StepResult48from .containers.runtime import LocalDockerProvider, UVProvider49from .utils import convert_to_ws_url50 51if TYPE_CHECKING:52 from websockets.asyncio.client import ClientConnection53 54 from .containers.runtime import ContainerProvider, RuntimeProvider55 from .sync_client import SyncEnvClient56 57from websockets.asyncio.client import connect as ws_connect58from websockets.protocol import State59 60ActT = TypeVar("ActT")61ObsT = TypeVar("ObsT")62EnvClientT = TypeVar("EnvClientT", bound="EnvClient")63ResultT = TypeVar("ResultT")64 65_VALID_CLIENT_MODES = ("simulation", "production")66 67 68class _AutoAsyncResult(Generic[ResultT]):69 """Awaitable result that can also be resolved by synchronous access."""70 71 def __init__(72 self,73 client: "EnvClient[Any, Any, Any]",74 coro_factory: Callable[[], Any],75 ):76 self._client = client77 self._coro_factory = coro_factory78 self._resolved = False79 self._result: ResultT | None = None80 81 def __await__(self):82 async def _await_result():83 self._client._claim_execution_mode("async")84 return await self._coro_factory()85 86 return _await_result().__await__()87 88 def _sync_result(self) -> ResultT:89 if not self._resolved:90 self._result = self._client._run_sync(self._coro_factory)91 self._resolved = True92 return self._result93 94 def __getattr__(self, name: str) -> Any:95 return getattr(self._sync_result(), name)96 97 def __bool__(self) -> bool:98 return bool(self._sync_result())99 100 def __repr__(self) -> str:101 if self._resolved:102 return repr(self._result)103 return f"<{type(self).__name__} pending>"104 105 106class _BootstrapResult(Coroutine, Generic[EnvClientT]):107 """Bootstrap handle returned by the client factory methods.108 109 Returned by [`~openenv.core.EnvClient.from_docker_image`] and110 [`~openenv.core.EnvClient.from_env`]. Resolving the handle is what actually111 starts the container / Space and connects the WebSocket, so a single factory112 call serves both execution modes:113 114 - `await handle` connects on the running event loop and returns the connected115 async client (unchanged async behavior).116 - `handle.sync()` connects on the sync background loop and returns a117 `SyncEnvClient`, mirroring the instance-level [`~openenv.core.EnvClient.sync`].118 119 The handle is a full coroutine (it implements `send` / `throw` / `close`), so120 it is a drop-in for the previous `async def` factories: `asyncio.run(...)`,121 `run_async_safely(...)`, and bare `await` all accept it, in addition to the122 new `.sync()` chain. Bootstrap is lazy — the underlying coroutine (and the123 container/Space start) is created only when the handle is driven or `.sync()`124 is called.125 """126 127 def __init__(self, bootstrap: Callable[[], EnvClientT]):128 self._bootstrap = bootstrap129 self._coro: Optional[Coroutine[Any, Any, EnvClientT]] = None130 self._used = False131 132 def _consume(self) -> EnvClientT:133 """Run the bootstrap exactly once; a second resolve is a programming error.134 135 Mirrors native coroutine semantics (a coroutine cannot be awaited twice)136 so that a handle re-driven via a second `.sync()`, or `await` followed by137 `.sync()`, raises instead of silently starting a second container/Space.138 """139 if self._used:140 raise RuntimeError(141 "This bootstrap handle has already been resolved; call the "142 "factory again to start a new environment."143 )144 self._used = True145 return self._bootstrap()146 147 async def _resolve_async(self) -> EnvClientT:148 client = self._consume()149 await client.connect()150 return client151 152 def _ensure_coro(self) -> "Coroutine[Any, Any, EnvClientT]":153 if self._coro is None:154 self._coro = self._resolve_async()155 return self._coro156 157 def __await__(self):158 return self._ensure_coro().__await__()159 160 def send(self, value: Any) -> Any:161 return self._ensure_coro().send(value)162 163 def throw(self, *args: Any, **kwargs: Any) -> Any:164 return self._ensure_coro().throw(*args, **kwargs)165 166 def close(self) -> None:167 if self._coro is not None:168 self._coro.close()169 170 def sync(self) -> "SyncEnvClient":171 client = self._consume()172 try:173 client._run_sync(client._connect_async)174 except Exception:175 # _consume() already started the provider. On the async path176 # _connect_async releases it, but here the failure may be the sync177 # loop setup itself (before _connect_async runs), so stop the178 # provider directly rather than routing through the broken loop.179 client._stop_provider_best_effort()180 raise181 return client.sync()182 183 def __repr__(self) -> str:184 return f"<{type(self).__name__} pending (await, run, or call .sync())>"185 186 187def _normalize_mode(mode: Optional[str]) -> str:188 """Resolve and validate the client communication mode."""189 raw_mode = (190 os.environ.get("OPENENV_CLIENT_MODE", "simulation") if mode is None else mode191 )192 normalized_mode = raw_mode.lower()193 if normalized_mode not in _VALID_CLIENT_MODES:194 raise ValueError(195 f"Invalid mode: '{normalized_mode}'. Must be 'simulation' or 'production'. "196 f"Set via constructor parameter or OPENENV_CLIENT_MODE environment variable."197 )198 return normalized_mode199 200 201def _is_localhost_ws_url(ws_url: str) -> bool:202 """Return True when the WebSocket URL targets the local loopback interface.203 204 The hostname is parsed from the URL so that only the actual host is matched.205 Substring matching is avoided because remote hosts such as206 ``my-localhost-proxy.example.com`` or ``127.0.0.1.example.com`` must not be207 treated as local.208 """209 hostname = urlsplit(ws_url).hostname210 if hostname is None:211 return False212 if hostname == "localhost":213 return True214 try:215 return ipaddress.ip_address(hostname).is_loopback216 except ValueError:217 return False218 219 220def _required_start_container_parameters(provider: Any) -> list[str]:221 """Return required arguments for a bound provider.start_container()."""222 try:223 signature = inspect.signature(provider.start_container)224 except (TypeError, ValueError):225 return []226 return [227 name228 for name, parameter in signature.parameters.items()229 if parameter.default is inspect.Parameter.empty230 and parameter.kind231 in (232 inspect.Parameter.POSITIONAL_ONLY,233 inspect.Parameter.POSITIONAL_OR_KEYWORD,234 inspect.Parameter.KEYWORD_ONLY,235 )236 ]237 238 239async def _best_effort_close(ws: ClientConnection) -> None:240 """Close a socket without letting a slow handshake or failure propagate.241 242 Scheduled as a background task (never awaited directly) so a dropped243 socket's close handshake can't hold up the caller that triggered the244 drop -- see `EnvClient._receive()`. `EnvClient._close_async()` awaits245 any still-running instance of this task before it returns, so a real246 `close()` call does wait for the handshake; only the caller that247 happened to trigger the drop is spared. `CancelledError` is caught248 alongside `Exception` because that awaiting is exactly what would249 otherwise cancel this mid-handshake if a caller called `close()`250 concurrently.251 """252 try:253 await ws.close()254 except (Exception, asyncio.CancelledError):255 pass # Best effort256 257 258async def _best_effort_graceful_close(ws: ClientConnection) -> None:259 """Notify the server, then close a socket without propagating failures."""260 try:261 await ws.send(json.dumps({"type": "close"}))262 except (Exception, asyncio.CancelledError):263 pass # Best effort264 await _best_effort_close(ws)265 266 267class EnvClient(ABC, Generic[ActT, ObsT, StateT]):268 """269 Async environment client for persistent sessions.270 271 This client maintains a persistent WebSocket connection to an environment272 server, enabling efficient multi-step interactions. Each client instance273 corresponds to a dedicated environment session on the server.274 275 The client is async by default. For synchronous usage, use the `.sync()`276 method to get a `SyncEnvClient` wrapper.277 278 Features:279 - Lower latency for sequential interactions280 - Session state is maintained server-side281 - Better suited for long-running episodes282 - Async by default for modern Python async/await patterns283 284 Examples:285 286 Async usage:287 288 ```python289 from envs.coding_env.client import CodingEnv290 291 async with CodingEnv(base_url="ws://localhost:8000") as env:292 result = await env.reset(seed=42)293 while not result.done:294 action = agent.predict(result.observation)295 result = await env.step(action)296 ```297 298 Sync usage via `.sync()` wrapper:299 300 ```python301 env = CodingEnv(base_url="ws://localhost:8000").sync()302 with env:303 result = env.reset(seed=42)304 result = env.step(action)305 ```306 """307 308 def __init__(309 self,310 base_url: Optional[str] = None,311 connect_timeout_s: float = 10.0,312 message_timeout_s: float = 60.0,313 max_message_size_mb: float = 100.0,314 websocket_ping_interval_s: Optional[float] = 20.0,315 websocket_ping_timeout_s: Optional[float] = 20.0,316 provider: Optional["ContainerProvider | RuntimeProvider"] = None,317 mode: Optional[str] = None,318 ):319 """320 Initialize environment client.321 322 Args:323 base_url (`str`, *optional*):324 Base URL of the environment server (http:// or ws://). Will be converted to325 ws:// if http:// is provided. May be omitted when the provider326 has enough constructor state to start itself.327 connect_timeout_s (`float`, *optional*, defaults to `10.0`):328 Timeout for establishing WebSocket connection.329 message_timeout_s (`float`, *optional*, defaults to `60.0`):330 Timeout for receiving responses to messages.331 max_message_size_mb (`float`, *optional*, defaults to `100.0`):332 Maximum WebSocket message size in megabytes. Default 100MB to handle large333 observations (screenshots, DOM, etc.).334 websocket_ping_interval_s (`float` or `None`, *optional*, defaults to `20.0`):335 WebSocket keepalive ping interval. Pass `None` to disable.336 websocket_ping_timeout_s (`float` or `None`, *optional*, defaults to `20.0`):337 WebSocket keepalive pong timeout. Pass `None` to disable.338 provider (`ContainerProvider` or `RuntimeProvider`, *optional*):339 Container/runtime provider for lifecycle management.340 mode (`str`, *optional*):341 Communication mode: `'simulation'` for Gym-style API (default) or342 `'production'` for MCP JSON-RPC protocol. Can also be set via the343 `OPENENV_CLIENT_MODE` environment variable. Constructor parameter takes344 precedence over environment variable. Case-insensitive.345 """346 if base_url is None and provider is None:347 raise ValueError("EnvClient requires either base_url or provider.")348 349 # Store mode (use object.__setattr__ to bypass immutability)350 object.__setattr__(self, "_mode", _normalize_mode(mode))351 352 self._base_url: Optional[str] = None353 self._ws_url: Optional[str] = None354 self._connect_timeout = connect_timeout_s355 self._message_timeout = message_timeout_s356 self._max_message_size = int(357 max_message_size_mb * 1024 * 1024358 ) # Convert MB to bytes359 self._websocket_ping_interval_s = websocket_ping_interval_s360 self._websocket_ping_timeout_s = websocket_ping_timeout_s361 self._provider = provider362 self._provider_stopped = False363 self._provider_cleanup_pending = False364 self._start_provider_on_connect = base_url is None365 self._child_clients: list[EnvClient[Any, Any, Any]] = []366 self._ws: Optional[ClientConnection] = None367 self._execution_mode: Optional[str] = None368 self._sync_client: Optional["SyncEnvClient[ActT, ObsT, StateT]"] = None369 self._ws_loop: Optional[asyncio.AbstractEventLoop] = None370 # Strong references for fire-and-forget socket closes (see _receive),371 # so the task isn't garbage-collected mid-close. Discarded on done.372 self._pending_close_tasks: set[asyncio.Task] = set()373 if base_url is not None:374 self._set_base_url(base_url)375 376 @property377 def base_url(self) -> Optional[str]:378 """Public read-only URL of the environment server this client targets.379 380 Returns the normalized `http(s)`/`ws(s)` base URL (no trailing slash),381 or `None` when the client was created with a provider but has not yet382 started its container (the URL is assigned lazily on `connect()`).383 384 This is the public counterpart to the private `self._base_url`; the385 previously-public `base_url` attribute was dropped in the `core`386 refactor, leaving sync consumers with no way to read the URL back.387 """388 return self._base_url389 390 def _set_base_url(self, base_url: str) -> None:391 self._base_url = base_url.rstrip("/")392 ws_url = convert_to_ws_url(base_url)393 self._ws_url = f"{ws_url}/ws"394 395 def _start_provider_if_needed(self) -> None:396 # A missing URL does not mean the previous resource was released.397 # Retry its cleanup before start can overwrite the provider's handle.398 if self._provider_cleanup_pending:399 if not self._start_provider_on_connect:400 raise RuntimeError(401 "Provider cleanup is pending for this client with an existing "402 "base URL. Retry close() to finish cleanup, then create a new "403 "client instead of reconnecting to the old URL."404 )405 self._stop_provider()406 if self._ws_url is not None:407 return408 if self._provider is None:409 raise RuntimeError("EnvClient has no base URL or provider.")410 if hasattr(self._provider, "start_container"):411 required_parameters = _required_start_container_parameters(self._provider)412 if required_parameters:413 required = ", ".join(required_parameters)414 raise ValueError(415 f"{type(self._provider).__name__} does not support "416 "provider-owned startup because start_container() requires "417 f"{required}. Start the provider manually and pass base_url, "418 "or configure a provider with a constructor-owned image/source."419 )420 self._provider_stopped = False421 base_url = self._provider.start_container()422 self._provider.wait_for_ready(base_url)423 elif hasattr(self._provider, "start"):424 self._provider_stopped = False425 base_url = self._provider.start()426 self._provider.wait_for_ready()427 else:428 raise TypeError("provider must define start_container() or start().")429 self._set_base_url(base_url)430 431 def _create_session_client(self) -> "EnvClient[Any, Any, Any]":432 # Match _start_provider_if_needed's startup condition. A provider with433 # an existing URL may already be serving the parent or other children.434 starts_provider_here = self._provider is not None and self._ws_url is None435 try:436 self._start_provider_if_needed()437 if self._base_url is None:438 raise RuntimeError("EnvClient has no base URL.")439 440 signature = inspect.signature(type(self))441 accepts_kwargs = any(442 parameter.kind == inspect.Parameter.VAR_KEYWORD443 for parameter in signature.parameters.values()444 )445 candidate_kwargs = {446 "base_url": self._base_url,447 "connect_timeout_s": self._connect_timeout,448 "message_timeout_s": self._message_timeout,449 "max_message_size_mb": self._max_message_size / (1024 * 1024),450 "websocket_ping_interval_s": self._websocket_ping_interval_s,451 "websocket_ping_timeout_s": self._websocket_ping_timeout_s,452 "mode": self._mode,453 }454 constructor_kwargs = {}455 for name, value in candidate_kwargs.items():456 if accepts_kwargs or name in signature.parameters:457 constructor_kwargs[name] = value458 459 return type(self)(**constructor_kwargs)460 except Exception:461 if starts_provider_here and not self._provider_cleanup_pending:462 self._stop_provider_best_effort()463 raise464 465 async def new_session(self) -> "EnvClient[Any, Any, Any]":466 """467 Create and connect a new session against the same environment server.468 469 Returns:470 `EnvClient`: A connected child client of the same concrete type.471 472 The child session is tracked by this parent and closed when the parent473 is closed. Server-side capacity still applies: when the server is at474 `MAX_CONCURRENT_ENVS`, opening the child WebSocket can fail and is475 surfaced as a connection error.476 """477 client = self._create_session_client()478 await client.connect()479 self._child_clients.append(client)480 return client481 482 def __setattr__(self, name: str, value: Any) -> None:483 """Prevent modification of _mode after initialization."""484 if name == "_mode" and hasattr(self, "_mode"):485 raise AttributeError("Cannot modify mode after initialization")486 super().__setattr__(name, value)487 488 def _claim_execution_mode(self, mode: str) -> None:489 """Lock the client to sync or async execution on first use."""490 if self._execution_mode is None:491 self._execution_mode = mode492 elif self._execution_mode != mode:493 raise RuntimeError(494 f"EnvClient is already being used in {self._execution_mode} mode. "495 "Create a separate client instance when mixing sync and async code."496 )497 498 def _run_sync(499 self,500 coro_factory: Callable[[], Any],501 *,502 allow_async_handoff: bool = False,503 ) -> Any:504 """Run an async operation through the sync wrapper."""505 if allow_async_handoff and self._execution_mode == "async":506 self._execution_mode = "sync"507 else:508 self._claim_execution_mode("sync")509 if self._sync_client is None:510 self._sync_client = self.sync()511 return self._sync_client._run(coro_factory())512 513 def _dispatch(self, coro_factory: Callable[[], Any]) -> Any:514 """Return an awaitable in async code and a concrete result in sync code."""515 try:516 running_loop = asyncio.get_running_loop()517 except RuntimeError:518 running_loop = None519 520 if self._execution_mode == "sync":521 sync_loop = (522 self._sync_client._loop if self._sync_client is not None else None523 )524 if running_loop is sync_loop:525 return coro_factory()526 return self._run_sync(coro_factory)527 if running_loop is not None:528 if self._execution_mode == "async":529 return coro_factory()530 return _AutoAsyncResult(self, coro_factory)531 532 return self._run_sync(coro_factory, allow_async_handoff=True)533 534 def connect(self) -> Any:535 return self._dispatch(self._connect_async)536 537 async def _connect_async(self) -> "EnvClient":538 """539 Establish WebSocket connection to the server.540 541 Returns:542 self for method chaining543 544 Raises:545 ConnectionError: If connection cannot be established546 """547 if self._ws is not None:548 if self._ws.state in (State.CLOSING, State.CLOSED):549 # Closed by the far end: a keepalive timeout, a tunnel dropping550 # the socket, the server restarting. The object is still not551 # None, so without this check it stays cached and every later552 # call raises `ConnectionClosed` for the rest of the process's553 # life. Only a demonstrably closed socket is dropped, so one554 # still CONNECTING is left alone.555 self._ws = None556 self._ws_loop = None557 elif self._ws_loop is asyncio.get_running_loop():558 return self559 # Connected from a different event loop than the one running560 # now -- e.g. `client = await Client.from_env(...)` inside561 # `asyncio.run(...)`, then `client.sync()` drives every later562 # call on `SyncEnvClient`'s own dedicated background loop. The563 # websocket object is bound to internals of the original loop,564 # which is typically already closed by the time we get here, so565 # it cannot be reused (or even cleanly closed) from this loop.566 # Drop the stale reference and reconnect fresh below rather than567 # silently no-op-ing onto a dead connection.568 self._ws = None569 self._ws_loop = None570 571 # A timed-out request drops its socket immediately but closes it in the572 # background so the timeout itself remains prompt. Wait for that close573 # before opening a replacement: the old server-side session continues574 # occupying a capacity slot until the close handshake finishes, and575 # many environments allow only one session.576 await self._drain_pending_close_tasks()577 578 try:579 self._start_provider_if_needed()580 except Exception:581 # A failed cleanup retry must propagate without attempting stop582 # again. For a new startup failure, preserve its original error.583 if not self._provider_cleanup_pending:584 with suppress(Exception):585 await self.close()586 raise587 588 assert self._ws_url is not None589 590 # Disable the proxy for localhost connections via the per-connection591 # `proxy` argument rather than mutating the process-global NO_PROXY592 # env var: concurrent connect() calls (e.g. asyncio.gather over many593 # env clients) would otherwise race on os.environ and leak state.594 connect_kwargs: Dict[str, Any] = {}595 if _is_localhost_ws_url(self._ws_url):596 connect_kwargs["proxy"] = None597 598 try:599 self._ws = await ws_connect(600 self._ws_url,601 open_timeout=self._connect_timeout,602 max_size=self._max_message_size,603 ping_interval=self._websocket_ping_interval_s,604 ping_timeout=self._websocket_ping_timeout_s,605 **connect_kwargs,606 )607 self._ws_loop = asyncio.get_running_loop()608 except Exception as e:609 await self.close()610 raise ConnectionError(f"Failed to connect to {self._ws_url}: {e}") from e611 612 return self613 614 def disconnect(self) -> Any:615 return self._dispatch(self._disconnect_async)616 617 def _schedule_socket_close(618 self, ws: ClientConnection, *, notify_server: bool = False619 ) -> asyncio.Task[None]:620 """Schedule and track a socket close until its handshake finishes."""621 close_coro = (622 _best_effort_graceful_close(ws) if notify_server else _best_effort_close(ws)623 )624 close_task = asyncio.create_task(close_coro)625 self._pending_close_tasks.add(close_task)626 close_task.add_done_callback(self._pending_close_tasks.discard)627 return close_task628 629 async def _disconnect_async(self) -> None:630 """Close the WebSocket connection."""631 if self._ws is not None:632 ws = self._ws633 ws_loop = self._ws_loop634 # Detach first so cancellation during the close handshake cannot635 # leave a stale socket cached for a later operation.636 self._ws = None637 self._ws_loop = None638 same_loop = ws_loop is asyncio.get_running_loop()639 if same_loop:640 # Keep ownership of the handshake if this caller is cancelled.641 # A later close/reconnect drains the task before tearing down642 # the loop or opening a replacement server session.643 close_task = self._schedule_socket_close(ws, notify_server=True)644 await asyncio.shield(close_task)645 646 async def _drain_pending_close_tasks(self) -> None:647 """Wait for background socket closes owned by the current event loop.648 649 Shielding keeps cancellation of the caller from cancelling the close650 tasks themselves. This matters both before reconnecting, when the old651 server session must release its capacity slot, and during explicit652 client shutdown.653 """654 loop = asyncio.get_running_loop()655 tasks = [656 task657 for task in tuple(self._pending_close_tasks)658 if not task.done() and task.get_loop() is loop659 ]660 if tasks:661 await asyncio.shield(asyncio.gather(*tasks, return_exceptions=True))662 663 async def _ensure_connected(self) -> None:664 """Ensure WebSocket connection is established on the current loop.665 666 Always delegates to `_connect_async()` rather than pre-checking667 `self._ws is None`: `_connect_async()` itself is the one that knows668 whether an existing `_ws` is reusable (same event loop) or stale (a669 different one, e.g. from a prior `from_env()` call now being driven670 through `.sync()`'s own loop).671 """672 await self._connect_async()673 674 async def _send(self, message: Dict[str, Any]) -> None:675 """Send a message over the WebSocket."""676 await self._ensure_connected()677 assert self._ws is not None678 await self._ws.send(json.dumps(message))679 680 async def _receive(self) -> Dict[str, Any]:681 """Receive the response on the same WebSocket that carried the request.682 683 Reconnecting here would wait on a fresh socket where the request was684 never sent. The next complete operation may reconnect in `_send()`.685 """686 assert self._ws is not None687 try:688 raw = await asyncio.wait_for(self._ws.recv(), timeout=self._message_timeout)689 except (asyncio.TimeoutError, asyncio.CancelledError):690 # The server may still write the response for this request to the691 # socket after we give up waiting on it. If we left the socket692 # open, the next call's `_receive()` would read that stale frame693 # and silently pair it with an unrelated request -- and every694 # response after that would be shifted by one for the life of695 # the connection. Drop the socket so the next `_send()` opens a696 # fresh one instead of reusing a desynced one.697 #698 # An outer `asyncio.wait_for` around the whole call (e.g. a699 # caller-imposed deadline on step()) cancels this await with700 # CancelledError, not our own message_timeout's TimeoutError --701 # same desync, different exception, so both are caught here.702 ws = self._ws703 self._ws = None704 self._ws_loop = None705 # Fire-and-forget: `close()` waits out the library's close706 # handshake (`close_timeout`, 10s by default) if the server is707 # slow to ack, and *awaiting* that here would hold up this708 # exception -- so an outer `asyncio.wait_for(..., timeout=50ms)`709 # would actually block for up to 10s before its deadline was710 # honored. Scheduling it lets the exception propagate711 # immediately while the close still happens in the background.712 self._schedule_socket_close(ws)713 raise714 return json.loads(raw)715 716 async def _send_and_receive(self, message: Dict[str, Any]) -> Dict[str, Any]:717 """Send a message and wait for response."""718 await self._send(message)719 response = await self._receive()720 721 # Check for error response722 if response.get("type") == "error":723 error_data = response.get("data", {})724 raise RuntimeError(725 f"Server error: {error_data.get('message', 'Unknown error')} "726 f"(code: {error_data.get('code', 'UNKNOWN')})"727 )728 729 return response730 731 @classmethod732 def _bootstrap_container(733 cls: Type[EnvClientT],734 image: str,735 provider: Optional["ContainerProvider"] = None,736 **kwargs: Any,737 ) -> EnvClientT:738 """Start a Docker container and build an *unconnected* client for it.739 740 Invoked lazily by the `from_docker_image` bootstrap handle when it is741 awaited or `.sync()`'d: the container start / readiness wait is plain742 blocking code, so the only thing that differs between the async and sync743 resolution paths is how the WebSocket is connected (awaited vs. run sync).744 """745 if provider is None:746 provider = LocalDockerProvider()747 748 try:749 # Start container, wait for readiness, then build the client.750 base_url = provider.start_container(image, **kwargs)751 provider.wait_for_ready(base_url)752 return cls(base_url=base_url, provider=provider)753 except Exception:754 # No EnvClient exists yet for the caller to close(), so release the755 # container here if start / readiness / construction fails.756 provider.stop_container()757 raise758 759 @classmethod760 def from_docker_image(761 cls: Type[EnvClientT],762 image: str,763 provider: Optional["ContainerProvider"] = None,764 **kwargs: Any,765 ) -> "_BootstrapResult[EnvClientT]":766 """767 Create an environment client by spinning up a Docker container.768 769 Returns a bootstrap handle so the same call works from both async and770 synchronous code: `await` it for the connected async client, or chain771 `.sync()` for a connected `SyncEnvClient` (e.g. a TRL GRPO rollout loop772 that cannot `await`). Bootstrap is lazy — the container starts when the773 handle is resolved.774 775 Args:776 image (`str`):777 Docker image name to run (e.g., `"coding-env:latest"`).778 provider (`ContainerProvider`, *optional*):779 Container provider to use. Defaults to `LocalDockerProvider`.780 **kwargs:781 Additional arguments to pass to `provider.start_container()`.782 783 Returns:784 `_BootstrapResult`: `await` for a connected async client, or call785 `.sync()` for a connected `SyncEnvClient`.786 787 Examples:788 789 ```python790 # Async791 env = await MyEnv.from_docker_image("coding-env:latest")792 793 # Sync794 env = MyEnv.from_docker_image("coding-env:latest").sync()795 result = env.reset()796 ```797 """798 return _BootstrapResult(799 lambda: cls._bootstrap_container(image, provider, **kwargs)800 )801 802 @classmethod803 def _bootstrap_env(804 cls: Type[EnvClientT],805 repo_id: str,806 *,807 use_docker: bool = True,808 provider: Optional["ContainerProvider | RuntimeProvider"] = None,809 **provider_kwargs: Any,810 ) -> EnvClientT:811 """Start a HF Space (docker or uv) and build an *unconnected* client.812 813 Invoked lazily by the `from_env` bootstrap handle; see814 `_bootstrap_container` for the same split rationale. On a startup815 failure the spawned process (and, for a git+ `project_path`, the temp816 clone directory) is released before re-raising; a later *connection*817 failure is cleaned up by `_connect_async` via `close()`.818 """819 # Extract start args that apply to both providers820 start_args = {}821 for key in ("port", "env_vars", "workers"):822 if key in provider_kwargs:823 start_args[key] = provider_kwargs.pop(key)824 825 if use_docker:826 # Docker mode: pull from HF registry827 docker_provider = provider or LocalDockerProvider()828 tag = provider_kwargs.pop("tag", "latest")829 image = f"registry.hf.space/{repo_id.replace('/', '-')}:{tag}"830 try:831 base_url = docker_provider.start_container(832 image, **start_args, **provider_kwargs833 )834 docker_provider.wait_for_ready(base_url)835 return cls(base_url=base_url, provider=docker_provider)836 except Exception:837 # No EnvClient exists yet for the caller to close(), so release838 # the container here if start / readiness / construction fails.839 docker_provider.stop_container()840 raise841 else:842 # UV mode: clone and run with uv843 if provider is None:844 uv_kwargs = dict(provider_kwargs)845 project_path = uv_kwargs.pop("project_path", None)846 if project_path is None:847 project_path = f"git+https://huggingface.co/spaces/{repo_id}"848 849 provider = UVProvider(project_path=project_path, **uv_kwargs)850 else:851 if provider_kwargs:852 raise ValueError(853 "provider_kwargs cannot be used when supplying a provider instance"854 )855 856 try:857 context_timeout_s = getattr(provider, "context_timeout_s", None)858 deadline = (859 time.monotonic() + context_timeout_s860 if context_timeout_s is not None861 else None862 )863 base_url = provider.start(**start_args)864 if deadline is None:865 provider.wait_for_ready()866 else:867 provider.wait_for_ready(868 timeout_s=max(0.0, deadline - time.monotonic())869 )870 return cls(base_url=base_url, provider=provider)871 except Exception:872 # No EnvClient exists yet for the caller to close(), so this is873 # the only chance to release the spawned process and (for a874 # git+ project_path) the temp clone directory. Covers start,875 # readiness, and client construction (e.g. an invalid mode).876 provider.stop()877 raise878 879 @classmethod880 def from_env(881 cls: Type[EnvClientT],882 repo_id: str,883 *,884 use_docker: bool = True,885 provider: Optional["ContainerProvider | RuntimeProvider"] = None,886 **provider_kwargs: Any,887 ) -> "_BootstrapResult[EnvClientT]":888 """889 Create a client from a Hugging Face Space.890 891 Returns a bootstrap handle: `await` it for the connected async client, or892 chain `.sync()` for a connected `SyncEnvClient`. Bootstrap is lazy — the893 Space starts when the handle is resolved.894 895 Args:896 repo_id (`str`):897 Hugging Face space identifier `{org}/{space}`.898 use_docker (`bool`, *optional*, defaults to `True`):899 When `True`, pull from the HF registry and launch via `LocalDockerProvider`.900 When `False`, run the space locally with `UVProvider`.901 provider (`ContainerProvider` or `RuntimeProvider`, *optional*):902 Provider instance to reuse. Must be a `ContainerProvider` when903 `use_docker=True` and a `RuntimeProvider` otherwise.904 **provider_kwargs:905 Additional keyword arguments forwarded to either the container provider's906 `start_container` (docker) or to the `UVProvider` constructor/start (uv).907 When `use_docker=False`, the `project_path` argument can be used to override908 the default git URL (`git+https://huggingface.co/spaces/{repo_id}`).909 910 Returns:911 `_BootstrapResult`: `await` for a connected async client, or call912 `.sync()` for a connected `SyncEnvClient`.913 914 Examples:915 916 ```python917 # Async: pull and run from HF Docker registry918 env = await MyEnv.from_env("openenv/echo-env")919 920 # Sync: chain .sync()921 env = MyEnv.from_env("openenv/echo-env").sync()922 923 # Run locally with UV (clones the space)924 env = await MyEnv.from_env("openenv/echo-env", use_docker=False)925 926 # Run from a local checkout927 env = await MyEnv.from_env(928 "openenv/echo-env",929 use_docker=False,930 project_path="/path/to/local/checkout"931 )932 ```933 """934 return _BootstrapResult(935 lambda: cls._bootstrap_env(936 repo_id,937 use_docker=use_docker,938 provider=provider,939 **provider_kwargs,940 )941 )942 943 @abstractmethod944 def _step_payload(self, action: ActT) -> Dict[str, Any]:945 """Convert an Action object to the JSON data expected by the env server."""946 raise NotImplementedError947 948 @abstractmethod949 def _parse_result(self, payload: Dict[str, Any]) -> StepResult[ObsT]:950 """Convert a JSON response from the env server to StepResult[ObsT]."""951 raise NotImplementedError952 953 @abstractmethod954 def _parse_state(self, payload: Dict[str, Any]) -> StateT:955 """Convert a JSON response from the state endpoint to a State object."""956 raise NotImplementedError957 958 def reset(self, **kwargs: Any) -> Any:959 return self._dispatch(lambda: self._reset_async(**kwargs))960 961 async def _reset_async(self, **kwargs: Any) -> StepResult[ObsT]:962 """963 Reset the environment with optional parameters.964 965 Args:966 **kwargs:967 Optional parameters passed to the environment's reset method.968 969 Returns:970 StepResult containing initial observation971 """972 message = {973 "type": "reset",974 "data": kwargs,975 }976 response = await self._send_and_receive(message)977 return self._parse_result(response.get("data", {}))978 979 def step(self, action: ActT, **kwargs: Any) -> Any:980 return self._dispatch(lambda: self._step_async(action, **kwargs))981 982 async def _step_async(self, action: ActT, **kwargs: Any) -> StepResult[ObsT]:983 """984 Execute an action in the environment.985 986 Args:987 action:988 The action to execute.989 **kwargs:990 Optional parameters (currently ignored).991 992 Returns:993 StepResult containing observation, reward, and done status994 """995 message = {996 "type": "step",997 "data": self._step_payload(action),998 }999 response = await self._send_and_receive(message)1000 return self._parse_result(response.get("data", {}))1001 1002 def state(self) -> Any:1003 return self._dispatch(self._state_async)1004 1005 async def _state_async(self) -> StateT:1006 """1007 Get the current environment state from the server.1008 1009 Returns:1010 State object with environment state information1011 """1012 message = {"type": "state"}1013 response = await self._send_and_receive(message)1014 return self._parse_state(response.get("data", {}))1015 1016 def close(self) -> Any:1017 return self._dispatch(self._close_async)1018 1019 async def _close_child_clients(self) -> asyncio.CancelledError | None:1020 """Close every captured child while deferring caller cancellation."""1021 children = list(self._child_clients)1022 self._child_clients.clear()1023 close_tasks = []1024 deferred_cancellation: asyncio.CancelledError | None = None1025 1026 for child in children:1027 try:1028 close_tasks.append(asyncio.ensure_future(child.close()))1029 except asyncio.CancelledError as exc:1030 if deferred_cancellation is None:1031 deferred_cancellation = exc1032 except Exception:1033 pass # Best effort1034 1035 if close_tasks:1036 close_group = asyncio.gather(*close_tasks, return_exceptions=True)1037 while not close_group.done():1038 try:1039 await asyncio.shield(close_group)1040 except asyncio.CancelledError as exc:1041 if deferred_cancellation is None:1042 deferred_cancellation = exc1043 1044 for result in close_group.result():1045 if (1046 isinstance(result, asyncio.CancelledError)1047 and deferred_cancellation is None1048 ):1049 deferred_cancellation = result1050 1051 return deferred_cancellation1052 1053 async def _close_async(self) -> None:1054 """1055 Close the WebSocket connection and clean up resources.1056 1057 If this client was created via from_docker_image() or from_env(),1058 this will also stop and remove the associated container/process.1059 """1060 deferred_cancellation: asyncio.CancelledError | None = None1061 try:1062 try:1063 try:1064 child_cancellation = await self._close_child_clients()1065 except asyncio.CancelledError as exc:1066 child_cancellation = exc1067 if child_cancellation is not None:1068 deferred_cancellation = child_cancellation1069 1070 # A real close waits out backgrounded closes, while shielding1071 # their socket handshakes from cancellation.1072 try:1073 await self._drain_pending_close_tasks()1074 except asyncio.CancelledError as exc:1075 if deferred_cancellation is None:1076 deferred_cancellation = exc1077 finally:1078 # Run even when child or pending-close cleanup is cancelled. A1079 # client may already have reconnected, and that current socket1080 # must not remain cached or open during teardown.1081 try:1082 await self._disconnect_async()1083 except asyncio.CancelledError as exc:1084 if deferred_cancellation is None:1085 deferred_cancellation = exc1086 finally:1087 self._stop_provider()1088 1089 if deferred_cancellation is not None:1090 raise deferred_cancellation1091 1092 def _stop_provider(self) -> None:1093 """Stop the provider once, retaining it for retries and later startup."""1094 try:1095 provider = self._provider1096 if provider is None or self._provider_stopped:1097 return1098 1099 if hasattr(provider, "stop_container"):1100 stop = provider.stop_container1101 elif hasattr(provider, "stop"):1102 stop = provider.stop1103 else:1104 return1105 1106 self._provider_cleanup_pending = True1107 stop()1108 # Only a successful stop discharges our cleanup responsibility.1109 self._provider_cleanup_pending = False1110 self._provider_stopped = True1111 finally:1112 if self._start_provider_on_connect:1113 self._base_url = None1114 self._ws_url = None1115 1116 def _stop_provider_best_effort(self) -> None:1117 """Stop the underlying provider directly, ignoring any errors.1118 1119 Releases a started container/process when there is no connected client1120 to `close()` through the normal path — e.g. sync bootstrap setup fails1121 after the provider started but before the connection is established, so1122 routing cleanup through the (possibly broken) sync loop is not an option.1123 """1124 with suppress(Exception):1125 self._stop_provider()1126 1127 async def __aenter__(self) -> "EnvClient":1128 """Enter async context manager, ensuring connection is established."""1129 self._claim_execution_mode("async")1130 await self.connect()1131 return self1132 1133 async def __aexit__(self, exc_type, exc_val, exc_tb) -> None:1134 """Exit async context manager, closing connection."""1135 await self.close()1136 1137 def __enter__(self) -> "EnvClient":1138 """Enter sync context manager, ensuring connection is established."""1139 self._run_sync(self._connect_async)1140 return self1141 1142 def __exit__(self, exc_type, exc_val, exc_tb) -> None:1143 """Exit sync context manager, closing connection."""1144 self._run_sync(self._close_async)1145 1146 def sync(self) -> "SyncEnvClient":1147 """1148 Return a synchronous wrapper around this async client.1149 1150 Use this method when you need synchronous access to the environment1151 without async/await syntax. This is useful for:1152 - Integration with synchronous codebases1153 - Interactive/REPL usage1154 - Stopping async from "infecting" the call stack1155 1156 Returns:1157 SyncEnvClient wrapper that provides synchronous methods1158 1159 Examples:1160 1161 ```python1162 async_client = GenericEnvClient(base_url="http://localhost:8000")1163 sync_client = async_client.sync()1164 1165 with sync_client:1166 result = sync_client.reset()1167 result = sync_client.step({"code": "print('hello')"})1168 ```1169 """1170 from .sync_client import SyncEnvClient1171 1172 if self._sync_client is None:1173 self._sync_client = SyncEnvClient(self)1174 return self._sync_client1175 