Team Ai
Apppublic

openenv/atari_env

sourceHugging Faceupdated 6mo agoView on Hugging Face
3likes
client.py121 linesDownload Raw Back to root
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"""8Atari Environment Client.9 10This module provides the client for connecting to an Atari Environment server11via WebSocket for persistent sessions.12"""13 14from __future__ import annotations15 16from typing import Any, Dict, TYPE_CHECKING17 18from openenv.core.client_types import StepResult19from openenv.core.env_client import EnvClient20 21from .models import AtariAction, AtariObservation, AtariState22 23if TYPE_CHECKING:24    from openenv.core.containers.runtime import ContainerProvider25 26 27class AtariEnv(EnvClient[AtariAction, AtariObservation, AtariState]):28    """29    Client for Atari Environment.30 31    This client maintains a persistent WebSocket connection to the environment32    server, enabling efficient multi-step interactions with lower latency.33 34    Example:35        >>> # Connect to a running server36        >>> with AtariEnv(base_url="http://localhost:8000") as client:37        ...     result = client.reset()38        ...     print(result.observation.screen_shape)39        ...40        ...     result = client.step(AtariAction(action_id=2))  # UP41        ...     print(result.reward, result.done)42 43    Example with Docker:44        >>> # Automatically start container and connect45        >>> client = AtariEnv.from_docker_image("atari-env:latest")46        >>> try:47        ...     result = client.reset()48        ...     result = client.step(AtariAction(action_id=0))  # NOOP49        ... finally:50        ...     client.close()51    """52 53    def _step_payload(self, action: AtariAction) -> Dict[str, Any]:54        """55        Convert AtariAction to JSON payload for step request.56 57        Args:58            action: AtariAction instance.59 60        Returns:61            Dictionary representation suitable for JSON encoding.62        """63        return {64            "action_id": action.action_id,65            "game_name": action.game_name,66            "obs_type": action.obs_type,67            "full_action_space": action.full_action_space,68        }69 70    def _parse_result(self, payload: Dict[str, Any]) -> StepResult[AtariObservation]:71        """72        Parse server response into StepResult[AtariObservation].73 74        Args:75            payload: JSON response from server.76 77        Returns:78            StepResult with AtariObservation.79        """80        obs_data = payload.get("observation", {})81 82        observation = AtariObservation(83            screen=obs_data.get("screen", []),84            screen_shape=obs_data.get("screen_shape", []),85            legal_actions=obs_data.get("legal_actions", []),86            lives=obs_data.get("lives", 0),87            episode_frame_number=obs_data.get("episode_frame_number", 0),88            frame_number=obs_data.get("frame_number", 0),89            done=payload.get("done", False),90            reward=payload.get("reward"),91            metadata=obs_data.get("metadata", {}),92        )93 94        return StepResult(95            observation=observation,96            reward=payload.get("reward"),97            done=payload.get("done", False),98        )99 100    def _parse_state(self, payload: Dict[str, Any]) -> AtariState:101        """102        Parse server response into AtariState object.103 104        Args:105            payload: JSON response from /state endpoint.106 107        Returns:108            AtariState object with environment state information.109        """110        return AtariState(111            episode_id=payload.get("episode_id"),112            step_count=payload.get("step_count", 0),113            game_name=payload.get("game_name", "unknown"),114            obs_type=payload.get("obs_type", "rgb"),115            full_action_space=payload.get("full_action_space", False),116            mode=payload.get("mode"),117            difficulty=payload.get("difficulty"),118            repeat_action_probability=payload.get("repeat_action_probability", 0.0),119            frameskip=payload.get("frameskip", 4),120        )121