SpaceFactory / env.py
Sahil Tailor
Initial Commit
92d87c0
Raw
History Blame Contribute Delete
3.63 kB
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 ──────────────────────────────────────────────────────────
@staticmethod
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)