Spaces:
Sleeping
Sleeping
| from __future__ import annotations | |
| from typing import Any, Dict, Optional, Tuple | |
| from .factory import ManufacturingFactoryEnv | |
| from .models import ( | |
| ManufacturingAction, | |
| ManufacturingEnvState, | |
| ManufacturingObservation, | |
| ManufacturingReward, | |
| ) | |
| from .tasks import TASK_REGISTRY, Task | |
| class ManufacturingTaskEnv: | |
| """Canonical user-facing API for the Space Manufacturing RL environment.""" | |
| def __init__( | |
| self, | |
| task_name: str = "easy", | |
| num_platforms: Optional[int] = None, | |
| max_steps: Optional[int] = None, | |
| ) -> None: | |
| self._task_name = task_name | |
| self._num_platforms_override = num_platforms | |
| self._max_steps_override = max_steps | |
| self._factory: Optional[ManufacturingFactoryEnv] = None | |
| self._task: Optional[Task] = None | |
| self.set_task(task_name) | |
| # ββ Class methods ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def list_tasks() -> Dict[str, str]: | |
| return {name: task.description for name, task in TASK_REGISTRY.items()} | |
| # ββ Setup ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def set_task(self, task_name: str) -> None: | |
| if task_name not in TASK_REGISTRY: | |
| raise ValueError( | |
| f"Unknown task '{task_name}'. Available: {list(TASK_REGISTRY.keys())}" | |
| ) | |
| self._task_name = task_name | |
| self._task = TASK_REGISTRY[task_name] | |
| self._factory = ManufacturingFactoryEnv( | |
| num_platforms=self._num_platforms_override or self._task.num_platforms, | |
| max_steps=self._max_steps_override or self._task.max_steps, | |
| seed=self._task.seed, | |
| pending_orders=list(self._task.pending_orders), | |
| delivery_windows=list(self._task.delivery_windows), | |
| solar_zones=dict(self._task.solar_zones), | |
| ) | |
| # ββ Core API βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def reset(self) -> ManufacturingObservation: | |
| assert self._factory is not None | |
| return self._factory.reset() | |
| def step( | |
| self, action: ManufacturingAction | |
| ) -> Tuple[ManufacturingObservation, ManufacturingReward, bool, Dict[str, Any]]: | |
| assert self._factory is not None | |
| return self._factory.step(action) | |
| def state(self) -> ManufacturingEnvState: | |
| assert self._factory is not None | |
| f = self._factory | |
| return ManufacturingEnvState( | |
| episode_id=f.episode_id, | |
| task_name=self._task_name, | |
| step_count=f.step_count, | |
| max_steps=f.max_steps, | |
| seed=f.seed, | |
| done=f.done, | |
| total_reward=f.total_reward, | |
| metrics=dict(f.metrics), | |
| platforms=list(f.platforms), | |
| delivery_windows=list(f.delivery_windows), | |
| solar_conditions=dict(f._solar_zones), | |
| pending_orders=list(f.pending_orders), | |
| ) | |
| # ββ Internal helper (for server compatibility) βββββββββββββββββββββββββββββ | |
| def _observation_from_dict(self, d: Dict[str, Any]) -> ManufacturingObservation: | |
| return ManufacturingObservation.model_validate(d) | |