| 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 |
|
|
| 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: |
| |
| 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) |
|
|