Spaces:
Sleeping
Sleeping
| # Copyright (c) Meta Platforms, Inc. and affiliates. | |
| # All rights reserved. | |
| # | |
| # This source code is licensed under the BSD-style license found in the | |
| # LICENSE file in the root directory of this source tree. | |
| """Typed models for the satellite environment.""" | |
| import json | |
| from typing import Any, Dict, List, Literal, Tuple | |
| import pydantic | |
| BaseModel = pydantic.BaseModel | |
| Field = pydantic.Field | |
| model_validator = getattr(pydantic, "model_validator", None) | |
| field_validator = getattr(pydantic, "field_validator", None) | |
| root_validator = getattr(pydantic, "root_validator", None) | |
| validator = getattr(pydantic, "validator", None) | |
| try: | |
| from openenv.core.env_server.types import Action, Observation | |
| except Exception: # pragma: no cover - local simulator fallback | |
| Action = BaseModel | |
| Observation = BaseModel | |
| class SatelliteState(BaseModel): | |
| """State of an individual satellite in the constellation.""" | |
| id: int = Field(..., description="Satellite ID") | |
| position: Tuple[float, float, float] = Field( | |
| ..., description="Position (x, y, z) in orbit" | |
| ) | |
| battery: float = Field(..., description="Battery level 0-100") | |
| storage: float = Field(..., description="Storage used percentage 0-100") | |
| last_action: str = Field(..., description="Last action taken") | |
| class SatelliteAction(Action): | |
| """Action for satellite constellation - commands for each satellite.""" | |
| def _normalize_payload(value: Any) -> Any: | |
| """Accept dicts, JSON strings, and common wrapper shapes from HTTP UIs.""" | |
| if isinstance(value, str): | |
| try: | |
| value = json.loads(value) | |
| except json.JSONDecodeError: | |
| return value | |
| if not isinstance(value, dict): | |
| return value | |
| if "action" in value and "satellite_actions" not in value: | |
| nested_action = value["action"] | |
| if isinstance(nested_action, str): | |
| try: | |
| nested_action = json.loads(nested_action) | |
| except json.JSONDecodeError: | |
| nested_action = value["action"] | |
| if isinstance(nested_action, dict): | |
| value = nested_action | |
| if "satellite_actions" not in value: | |
| return {"satellite_actions": value} | |
| satellite_actions = value.get("satellite_actions") | |
| if isinstance(satellite_actions, str): | |
| try: | |
| satellite_actions = json.loads(satellite_actions) | |
| except json.JSONDecodeError: | |
| return value | |
| if ( | |
| isinstance(satellite_actions, dict) | |
| and "satellite_actions" in satellite_actions | |
| and len(satellite_actions) == 1 | |
| ): | |
| satellite_actions = satellite_actions["satellite_actions"] | |
| if isinstance(satellite_actions, dict): | |
| value = {**value, "satellite_actions": satellite_actions} | |
| return value | |
| def _normalize_satellite_actions_field(value: Any) -> Any: | |
| """Handle direct field validation when wrappers are passed to the field itself.""" | |
| if isinstance(value, str): | |
| try: | |
| value = json.loads(value) | |
| except json.JSONDecodeError: | |
| return value | |
| if not isinstance(value, dict): | |
| return value | |
| if "action" in value: | |
| nested_action = value["action"] | |
| if isinstance(nested_action, str): | |
| try: | |
| nested_action = json.loads(nested_action) | |
| except json.JSONDecodeError: | |
| nested_action = value["action"] | |
| if isinstance(nested_action, dict): | |
| value = nested_action | |
| if "satellite_actions" in value: | |
| nested_actions = value["satellite_actions"] | |
| if isinstance(nested_actions, str): | |
| try: | |
| nested_actions = json.loads(nested_actions) | |
| except json.JSONDecodeError: | |
| return value | |
| if isinstance(nested_actions, dict): | |
| return nested_actions | |
| return value | |
| if model_validator is not None: | |
| def _coerce_payload(cls, value: Any) -> Any: | |
| return cls._normalize_payload(value) | |
| elif root_validator is not None: # pragma: no cover - pydantic v1 fallback | |
| def _coerce_payload(cls, values: Dict[str, Any]) -> Dict[str, Any]: | |
| return cls._normalize_payload(values) | |
| if field_validator is not None: | |
| def _coerce_satellite_actions_field(cls, value: Any) -> Any: | |
| return cls._normalize_satellite_actions_field(value) | |
| elif validator is not None: # pragma: no cover - pydantic v1 fallback | |
| def _coerce_satellite_actions_field(cls, value: Any) -> Any: | |
| return cls._normalize_satellite_actions_field(value) | |
| satellite_actions: Dict[int, Literal["capture", "downlink", "maintain", "idle"]] = Field( | |
| ..., | |
| description=( | |
| "Dictionary mapping satellite_id to action: " | |
| "'capture', 'downlink', 'maintain', or 'idle'" | |
| ), | |
| ) | |
| class SatelliteObservation(Observation): | |
| """Observation from the Satellite environment.""" | |
| satellites: List[SatelliteState] = Field(..., description="States of all satellites") | |
| time_step: int = Field(..., description="Current time step") | |
| ground_stations: List[Tuple[float, float]] = Field( | |
| ..., description="Ground station coordinates (lat, lon)" | |
| ) | |
| weather_conditions: Dict[str, float] = Field(..., description="Cloud cover by region") | |
| pending_tasks: List[Dict[str, Any]] = Field(..., description="List of pending tasks") | |
| total_reward: float = Field(default=0.0, description="Cumulative reward") | |
| done: bool = Field(default=False, description="Whether the episode has ended") | |
| reward: float = Field(default=0.0, description="Immediate reward from the latest step") | |
| metadata: Dict[str, Any] = Field( | |
| default_factory=dict, description="Additional step metadata" | |
| ) | |
| class SatelliteReward(BaseModel): | |
| """Typed reward model returned by the canonical environment API.""" | |
| value: float = Field(..., description="Scalar reward for the latest step") | |
| components: Dict[str, float] = Field( | |
| default_factory=dict, description="Reward component breakdown" | |
| ) | |
| class SatelliteEnvState(BaseModel): | |
| """Typed state snapshot returned by state().""" | |
| episode_id: str = Field(..., description="Unique episode identifier") | |
| task_name: str = Field(..., description="Active task preset") | |
| step_count: int = Field(..., description="Current environment step count") | |
| max_steps: int = Field(..., description="Maximum steps allowed in the episode") | |
| seed: int = Field(..., description="Deterministic seed for the current task setup") | |
| done: bool = Field(..., description="Whether the episode is finished") | |
| total_reward: float = Field(..., description="Accumulated reward") | |
| metrics: Dict[str, float] = Field( | |
| default_factory=dict, description="Deterministic episode metrics used for grading" | |
| ) | |
| satellites: List[SatelliteState] = Field(..., description="Current satellite states") | |
| ground_stations: List[Tuple[float, float]] = Field( | |
| ..., description="Ground station coordinates (lat, lon)" | |
| ) | |
| weather_conditions: Dict[str, float] = Field(..., description="Cloud cover by region") | |
| pending_tasks: List[Dict[str, Any]] = Field(..., description="Remaining visible tasks") | |