openenv/atari_env
3
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 