File size: 1,743 Bytes
7043cc6 aaf289c 7043cc6 aaf289c 7043cc6 6d193c2 aaf289c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 | try:
from openenv.core.client import EnvClient
except ImportError:
try:
from openenv.core.env_client import EnvClient
except ImportError:
from openenv.core import EnvClient
from typing import Dict, Any
from models import AdaptiveAction, AdaptiveObservation, AdaptiveState
try:
from openenv.core.models import StepResult
except ImportError:
pass # We will construct dictionaries if StepResult is missing in older versions
class AdaptiveWorldEnv(EnvClient[AdaptiveAction, AdaptiveObservation, AdaptiveState]):
"""
Client for the AdaptiveWorld environment.
Connects to the HF Space or local server running adaptive-world-env.
"""
def __init__(self, base_url: str = "http://localhost:7860"):
super().__init__(base_url=base_url)
def _step_payload(self, action: AdaptiveAction) -> Dict[str, Any]:
return action.model_dump()
def _parse_result(self, payload: Dict[str, Any]) -> Any:
try:
from openenv.core.models import StepResult
return StepResult(
observation=AdaptiveObservation(**payload.get("observation", {})),
reward=payload.get("reward", 0.0),
done=payload.get("done", False),
info=payload.get("info", {})
)
except ImportError:
# Fallback for parsing dict as object
class Result:
def __init__(self, **kwargs):
self.__dict__.update(kwargs)
self.observation = AdaptiveObservation(**kwargs.get("observation", {}))
return Result(**payload)
def _parse_state(self, payload: Dict[str, Any]) -> AdaptiveState:
return AdaptiveState(**payload)
|