SpaceFactory / models.py
Sahil Tailor
Initial Commit
92d87c0
Raw
History Blame Contribute Delete
2.51 kB
from __future__ import annotations
from typing import Any, Dict, List, Literal, Optional
try:
from openenv.core.env_server.types import Action, Observation
_BASE_ACTION = Action
_BASE_OBS = Observation
except ImportError:
from pydantic import BaseModel as _BASE_ACTION # type: ignore
from pydantic import BaseModel as _BASE_OBS # type: ignore
from pydantic import BaseModel, field_validator, model_validator
PlatformActionType = Literal["produce", "assemble", "deliver", "recharge"]
class PlatformState(BaseModel):
id: int
position: List[float]
energy: float
material_stock: float
component_stock: float
product_stock: int
last_action: Optional[str] = None
class DeliveryWindow(BaseModel):
order_id: int
product_type: str
deadline: int
reward_value: float = 8.0
class PendingOrder(BaseModel):
order_id: int
product_type: str
requires_assembly: bool = False
quantity: int = 1
class ManufacturingAction(_BASE_ACTION):
platform_actions: Dict[int, PlatformActionType]
@model_validator(mode="before")
@classmethod
def coerce_from_http(cls, values: Any) -> Any:
# Accept flat dict of string keys from HTTP JSON
if isinstance(values, dict) and "platform_actions" not in values:
# Maybe the whole dict IS the platform_actions mapping
if all(isinstance(k, (int, str)) for k in values):
return {"platform_actions": values}
return values
@field_validator("platform_actions", mode="before")
@classmethod
def coerce_keys(cls, v: Any) -> Any:
if isinstance(v, dict):
return {int(k): val for k, val in v.items()}
return v
class ManufacturingReward(BaseModel):
value: float
components: Dict[str, float] = {}
class ManufacturingObservation(_BASE_OBS):
platforms: List[PlatformState]
time_step: int
delivery_windows: List[DeliveryWindow]
solar_conditions: Dict[str, float]
pending_orders: List[PendingOrder]
total_reward: float = 0.0
reward: float = 0.0
done: bool = False
metadata: Dict[str, Any] = {}
class ManufacturingEnvState(BaseModel):
episode_id: str
task_name: str
step_count: int
max_steps: int
seed: int
done: bool
total_reward: float
metrics: Dict[str, Any]
platforms: List[PlatformState]
delivery_windows: List[DeliveryWindow]
solar_conditions: Dict[str, float]
pending_orders: List[PendingOrder]