Spaces:
Sleeping
Sleeping
| import asyncio | |
| import importlib.util | |
| import json | |
| import math | |
| import os | |
| import re | |
| import sys | |
| import textwrap | |
| import time | |
| from dataclasses import dataclass | |
| from pathlib import Path | |
| from typing import Any, Dict, List, Optional, Tuple | |
| from urllib.parse import urlparse | |
| from openai import OpenAI | |
| try: | |
| from dotenv import load_dotenv | |
| except Exception: # pragma: no cover - optional convenience dependency | |
| load_dotenv = None | |
| try: | |
| from satellite import EasyTask, HardTask, MediumTask, SatelliteAction, SatelliteTaskEnv, TaskGrader | |
| except ImportError: # pragma: no cover - script execution from repo root fallback | |
| package_root = Path(__file__).resolve().parent | |
| spec = importlib.util.spec_from_file_location( | |
| "satellite", | |
| package_root / "__init__.py", | |
| submodule_search_locations=[str(package_root)], | |
| ) | |
| if spec is None or spec.loader is None: | |
| raise | |
| satellite = importlib.util.module_from_spec(spec) | |
| sys.modules["satellite"] = satellite | |
| spec.loader.exec_module(satellite) | |
| EasyTask = satellite.EasyTask | |
| HardTask = satellite.HardTask | |
| MediumTask = satellite.MediumTask | |
| SatelliteAction = satellite.SatelliteAction | |
| SatelliteTaskEnv = satellite.SatelliteTaskEnv | |
| TaskGrader = satellite.TaskGrader | |
| if load_dotenv is not None: | |
| load_dotenv() | |
| def warn_once(key: str, message: str) -> None: | |
| if key in WARNINGS_EMITTED: | |
| return | |
| WARNINGS_EMITTED.add(key) | |
| print(f"[warn] {message}", file=sys.stderr) | |
| def read_float_env(name: str, default: float) -> float: | |
| raw = os.getenv(name) | |
| if raw is None: | |
| return default | |
| try: | |
| return float(raw) | |
| except (TypeError, ValueError): | |
| warn_once(f"env:{name}", f"Invalid {name}={raw!r}; using default {default}.") | |
| return default | |
| def read_int_env(name: str, default: int) -> int: | |
| raw = os.getenv(name) | |
| if raw is None: | |
| return default | |
| try: | |
| return int(raw) | |
| except (TypeError, ValueError): | |
| warn_once(f"env:{name}", f"Invalid {name}={raw!r}; using default {default}.") | |
| return default | |
| WARNINGS_EMITTED: set[str] = set() | |
| API_BASE_URL = os.environ["API_BASE_URL"] | |
| API_KEY = os.environ["API_KEY"] | |
| MODEL_NAME = os.getenv("MODEL_NAME") | |
| BASELINE_POLICY = os.getenv("BASELINE_POLICY", "openai").lower() | |
| TEMPERATURE = read_float_env("TEMPERATURE", 0.0) | |
| MAX_TOKENS = read_int_env("MAX_TOKENS", 300) | |
| REQUEST_DELAY = read_float_env("REQUEST_DELAY", 0.0) | |
| REQUEST_TIMEOUT = read_float_env("REQUEST_TIMEOUT", 30.0) | |
| DEBUG = os.getenv("DEBUG", "false").lower() == "true" | |
| FALLBACK_ACTION = "idle" | |
| TASK_ORDER = ["easy", "medium", "hard"] | |
| TASK_TYPES = { | |
| "easy": EasyTask, | |
| "medium": MediumTask, | |
| "hard": HardTask, | |
| } | |
| ACTION_PATTERN = re.compile(r"(capture|downlink|maintain|idle)", re.IGNORECASE) | |
| ACTION_PREFIX_RE = re.compile(r"^(action|next action)\s*[:\-]\s*", re.IGNORECASE) | |
| SYSTEM_PROMPT = textwrap.dedent( | |
| """ | |
| You are managing a real-world satellite constellation. | |
| Reply with exactly one JSON object mapping satellite ids to actions. | |
| Valid actions: | |
| - capture | |
| - downlink | |
| - maintain | |
| - idle | |
| Decision guidance: | |
| - prefer capture only when a visible image task exists and the satellite has battery/storage margin | |
| - prefer downlink when a visible ground station task exists, especially if storage is high | |
| - use maintain to recover low-battery satellites before they become risky | |
| - avoid invalid, repeated, or wasteful actions | |
| - keep the fleet healthy across the full episode, not just the current step | |
| Output format: | |
| {"0": "capture", "1": "idle"} | |
| Do not include explanations or any extra text outside the JSON object. | |
| """ | |
| ).strip() | |
| class TaskRunResult: | |
| task_name: str | |
| score: float | |
| total_reward: float | |
| steps: int | |
| done: bool | |
| metrics: Dict[str, float] | |
| def build_history_lines(history: List[str]) -> str: | |
| if not history: | |
| return "None" | |
| return "\n".join(history[-6:]) | |
| def satellite_geo(position: Tuple[float, float, float]) -> Tuple[float, float, float]: | |
| x, y, z = position | |
| radius = math.sqrt((x * x) + (y * y) + (z * z)) | |
| if radius <= 0: | |
| return 0.0, 0.0, 0.0 | |
| lat = math.degrees(math.asin(z / radius)) | |
| lon = math.degrees(math.atan2(y, x)) | |
| altitude = max(0.0, radius - 6371.0) | |
| return float(lat), float(lon), float(altitude) | |
| def visibility_radius_rad(altitude_km: float) -> float: | |
| earth_radius_km = 6371.0 | |
| alt = max(0.0, altitude_km) | |
| horizon = math.acos(min(1.0, earth_radius_km / (earth_radius_km + alt))) | |
| return max(math.radians(35.0), min(math.radians(120.0), horizon + math.radians(50.0))) | |
| def great_circle_distance_rad(lat1: float, lon1: float, lat2: float, lon2: float) -> float: | |
| lat1_rad = math.radians(lat1) | |
| lon1_rad = math.radians(lon1) | |
| lat2_rad = math.radians(lat2) | |
| lon2_rad = math.radians(lon2) | |
| d_lat = lat2_rad - lat1_rad | |
| d_lon = lon2_rad - lon1_rad | |
| a = ( | |
| math.sin(d_lat / 2.0) ** 2 | |
| + math.cos(lat1_rad) * math.cos(lat2_rad) * math.sin(d_lon / 2.0) ** 2 | |
| ) | |
| return 2.0 * math.asin(min(1.0, math.sqrt(a))) | |
| def normalize_pending_tasks(tasks: List[Dict[str, Any]]) -> List[Dict[str, Any]]: | |
| normalized = [] | |
| for task in tasks: | |
| normalized.append({key: value for key, value in task.items()}) | |
| return normalized | |
| def extract_capture_regions(observation: Dict[str, Any]) -> Dict[str, Tuple[float, float]]: | |
| regions = observation.get("capture_regions") | |
| if isinstance(regions, dict) and regions: | |
| return { | |
| str(name): (float(coords[0]), float(coords[1])) | |
| for name, coords in regions.items() | |
| if isinstance(coords, (list, tuple)) and len(coords) == 2 | |
| } | |
| return { | |
| "region1": (18.5, 73.9), | |
| "region2": (34.0, -117.0), | |
| "region3": (-22.8, -43.2), | |
| } | |
| def find_visible_capture_task( | |
| sat: Dict[str, Any], | |
| observation: Dict[str, Any], | |
| ) -> Optional[Dict[str, Any]]: | |
| capture_regions = extract_capture_regions(observation) | |
| sat_lat, sat_lon, sat_alt = satellite_geo(tuple(sat["position"])) | |
| max_distance = visibility_radius_rad(sat_alt) | |
| candidates: List[Tuple[float, str, Dict[str, Any]]] = [] | |
| weather = observation.get("weather_conditions", {}) | |
| for task in normalize_pending_tasks(observation.get("pending_tasks", [])): | |
| if task.get("type") != "image_capture": | |
| continue | |
| region = str(task.get("region", "")) | |
| if region not in capture_regions: | |
| continue | |
| reg_lat, reg_lon = capture_regions[region] | |
| distance = great_circle_distance_rad(sat_lat, sat_lon, reg_lat, reg_lon) | |
| if distance > max_distance: | |
| continue | |
| priority = float(task.get("priority", 1)) | |
| cloud = float(weather.get(region, 0.5)) | |
| score = (priority * 3.0) + ((1.0 - cloud) * 2.0) - distance | |
| candidates.append((score, str(task.get("id", "")), task)) | |
| if not candidates: | |
| return None | |
| candidates.sort(key=lambda item: (-item[0], item[1])) | |
| return candidates[0][2] | |
| def find_visible_downlink_task( | |
| sat: Dict[str, Any], | |
| observation: Dict[str, Any], | |
| ) -> Optional[Dict[str, Any]]: | |
| stations = observation.get("ground_stations", []) | |
| sat_lat, sat_lon, sat_alt = satellite_geo(tuple(sat["position"])) | |
| max_distance = visibility_radius_rad(sat_alt) | |
| candidates: List[Tuple[float, str, Dict[str, Any]]] = [] | |
| for task in normalize_pending_tasks(observation.get("pending_tasks", [])): | |
| if task.get("type") != "data_downlink": | |
| continue | |
| station_id = int(task.get("station", 0)) | |
| if station_id < 0 or station_id >= len(stations): | |
| continue | |
| gs_lat, gs_lon = stations[station_id] | |
| distance = great_circle_distance_rad(sat_lat, sat_lon, float(gs_lat), float(gs_lon)) | |
| if distance > max_distance: | |
| continue | |
| priority = float(task.get("priority", 1)) | |
| units_remaining = float(task.get("units_remaining", 20.0)) | |
| completion_bias = 0.75 if float(sat["storage"]) >= units_remaining else 0.0 | |
| score = (priority * 3.0) + completion_bias - distance | |
| candidates.append((score, str(task.get("id", "")), task)) | |
| if not candidates: | |
| return None | |
| candidates.sort(key=lambda item: (-item[0], item[1])) | |
| return candidates[0][2] | |
| def format_observation(task_name: str, observation: Dict[str, Any]) -> str: | |
| satellites_info = [] | |
| for sat in observation.get("satellites", []): | |
| sat_lat, sat_lon, sat_alt = satellite_geo(tuple(sat["position"])) | |
| capture_task = find_visible_capture_task(sat, observation) | |
| downlink_task = find_visible_downlink_task(sat, observation) | |
| satellites_info.append( | |
| ( | |
| f" Satellite {sat['id']}: battery={sat['battery']:.1f}, " | |
| f"storage={sat['storage']:.1f}, last={sat['last_action']}, " | |
| f"lat={sat_lat:.1f}, lon={sat_lon:.1f}, alt={sat_alt:.1f}km, " | |
| f"capture_visible={capture_task is not None}, " | |
| f"downlink_visible={downlink_task is not None}" | |
| ) | |
| ) | |
| tasks_info = [] | |
| for task in observation.get("pending_tasks", [])[:12]: | |
| descriptor = task["type"] | |
| if task["type"] == "image_capture": | |
| descriptor += f" region={task.get('region')}" | |
| if task["type"] == "data_downlink": | |
| descriptor += f" station={task.get('station')}" | |
| descriptor += f" units={task.get('units_remaining', 0)}" | |
| tasks_info.append(f" - {descriptor} priority={task.get('priority', 1)}") | |
| weather_info = ", ".join( | |
| f"{region}={cover:.0%}" | |
| for region, cover in observation.get("weather_conditions", {}).items() | |
| ) | |
| return textwrap.dedent( | |
| f""" | |
| Task: {task_name} | |
| Time Step: {observation.get('time_step', 0)} | |
| Total Reward: {observation.get('total_reward', 0.0):.2f} | |
| Weather: {weather_info or 'n/a'} | |
| Satellites: | |
| {chr(10).join(satellites_info) or ' None'} | |
| Pending Tasks: | |
| {chr(10).join(tasks_info) or ' - None'} | |
| """ | |
| ).strip() | |
| def build_user_prompt( | |
| task_name: str, | |
| step: int, | |
| observation: Dict[str, Any], | |
| history: List[str], | |
| total_reward: float, | |
| ) -> str: | |
| return textwrap.dedent( | |
| f""" | |
| Step: {step} | |
| Aggregate reward so far: {total_reward:+.2f} | |
| Current state: | |
| {format_observation(task_name, observation)} | |
| Previous steps: | |
| {build_history_lines(history)} | |
| Reply with exactly one JSON object. | |
| """ | |
| ).strip() | |
| def heuristic_action(observation: Dict[str, Any]) -> Dict[int, str]: | |
| actions: Dict[int, str] = {} | |
| pending_tasks = observation.get("pending_tasks", []) | |
| has_capture_task = any(task.get("type") == "image_capture" for task in pending_tasks) | |
| has_downlink_task = any(task.get("type") == "data_downlink" for task in pending_tasks) | |
| for sat in observation.get("satellites", []): | |
| sat_id = int(sat["id"]) | |
| battery = float(sat["battery"]) | |
| storage = float(sat["storage"]) | |
| visible_capture = find_visible_capture_task(sat, observation) | |
| visible_downlink = find_visible_downlink_task(sat, observation) | |
| if battery <= 12: | |
| actions[sat_id] = "maintain" | |
| continue | |
| if battery < 28 and not visible_downlink: | |
| actions[sat_id] = "maintain" | |
| continue | |
| if storage >= 85 and visible_downlink: | |
| actions[sat_id] = "downlink" | |
| continue | |
| if visible_capture and battery >= 25 and storage <= 80: | |
| actions[sat_id] = "capture" | |
| continue | |
| if visible_downlink and storage > 0: | |
| actions[sat_id] = "downlink" | |
| continue | |
| if battery < 45 and not has_capture_task: | |
| actions[sat_id] = "maintain" | |
| continue | |
| if battery < 35 and storage <= 5: | |
| actions[sat_id] = "maintain" | |
| continue | |
| if has_downlink_task and storage >= 50: | |
| actions[sat_id] = "idle" | |
| continue | |
| if has_capture_task and battery >= 30 and storage < 70: | |
| actions[sat_id] = "idle" | |
| continue | |
| actions[sat_id] = "idle" | |
| return actions | |
| def parse_model_action(response_text: str, observation: Dict[str, Any]) -> Dict[int, str]: | |
| if not response_text: | |
| return safe_heuristic_action(observation, "empty model response") | |
| cleaned = ACTION_PREFIX_RE.sub("", response_text.strip()) | |
| try: | |
| json_match = re.search(r"\{.*\}", cleaned, re.DOTALL) | |
| if json_match: | |
| parsed = json.loads(json_match.group(0)) | |
| actions: Dict[int, str] = {} | |
| valid_ids = {int(sat["id"]) for sat in observation.get("satellites", [])} | |
| for key, value in parsed.items(): | |
| sat_id = int(key) | |
| if sat_id not in valid_ids: | |
| continue | |
| action = str(value).strip().lower() | |
| if not ACTION_PATTERN.fullmatch(action): | |
| action = FALLBACK_ACTION | |
| actions[sat_id] = action | |
| if actions: | |
| fallback = heuristic_action(observation) | |
| for sat_id in valid_ids: | |
| actions.setdefault(sat_id, fallback.get(sat_id, FALLBACK_ACTION)) | |
| return actions | |
| except (json.JSONDecodeError, TypeError, ValueError): | |
| pass | |
| warn_once("parse:model-response", "Model response was not valid JSON; using heuristic actions.") | |
| return safe_heuristic_action(observation, "invalid model response") | |
| def observation_to_dict(observation: Any) -> Dict[str, Any]: | |
| return { | |
| "satellites": [ | |
| { | |
| "id": sat.id, | |
| "position": sat.position, | |
| "battery": sat.battery, | |
| "storage": sat.storage, | |
| "last_action": sat.last_action, | |
| } | |
| for sat in observation.satellites | |
| ], | |
| "time_step": observation.time_step, | |
| "ground_stations": observation.ground_stations, | |
| "weather_conditions": observation.weather_conditions, | |
| "pending_tasks": observation.pending_tasks, | |
| "total_reward": observation.total_reward, | |
| "done": observation.done, | |
| "reward": observation.reward, | |
| "metadata": observation.metadata, | |
| } | |
| def build_idle_actions(observation: Dict[str, Any]) -> Dict[int, str]: | |
| actions: Dict[int, str] = {} | |
| for sat in observation.get("satellites", []): | |
| try: | |
| actions[int(sat["id"])] = FALLBACK_ACTION | |
| except (KeyError, TypeError, ValueError): | |
| continue | |
| return actions | |
| def safe_heuristic_action(observation: Dict[str, Any], reason: str) -> Dict[int, str]: | |
| try: | |
| return heuristic_action(observation) | |
| except Exception as exc: # noqa: BLE001 | |
| warn_once( | |
| f"heuristic:{reason}", | |
| f"Heuristic fallback failed after {reason}: {exc}. Returning all-idle actions.", | |
| ) | |
| return build_idle_actions(observation) | |
| def validate_api_base_url(base_url: Optional[str]) -> Optional[str]: | |
| if not base_url: | |
| return None | |
| cleaned = base_url.strip().rstrip("/") | |
| parsed = urlparse(cleaned) | |
| if parsed.scheme not in {"http", "https"} or not parsed.netloc: | |
| warn_once( | |
| "config:api-base-url", | |
| f"Invalid API_BASE_URL={base_url!r}; falling back to heuristic policy.", | |
| ) | |
| return None | |
| return cleaned | |
| def extract_response_text(completion: Any) -> str: | |
| choices = getattr(completion, "choices", None) | |
| if not choices: | |
| return "" | |
| message = getattr(choices[0], "message", None) | |
| if message is None: | |
| return "" | |
| content = getattr(message, "content", "") | |
| if isinstance(content, str): | |
| return content | |
| if isinstance(content, list): | |
| parts: List[str] = [] | |
| for item in content: | |
| if isinstance(item, dict): | |
| text = item.get("text") | |
| else: | |
| text = getattr(item, "text", None) | |
| if text: | |
| parts.append(str(text)) | |
| return "\n".join(parts) | |
| return str(content or "") | |
| def build_client() -> OpenAI: | |
| return OpenAI( | |
| base_url=os.environ["API_BASE_URL"], | |
| api_key=os.environ["API_KEY"], | |
| timeout=REQUEST_TIMEOUT, | |
| ) | |
| def choose_actions( | |
| client: Optional[OpenAI], | |
| task_name: str, | |
| step: int, | |
| observation: Dict[str, Any], | |
| history: List[str], | |
| total_reward: float, | |
| ) -> Dict[int, str]: | |
| if client is None: | |
| raise RuntimeError("OpenAI client not initialized") | |
| user_prompt = build_user_prompt(task_name, step, observation, history, total_reward) | |
| try: | |
| completion = client.chat.completions.create( | |
| model=MODEL_NAME, | |
| messages=[ | |
| {"role": "system", "content": SYSTEM_PROMPT}, | |
| {"role": "user", "content": user_prompt}, | |
| ], | |
| temperature=TEMPERATURE, | |
| max_tokens=MAX_TOKENS, | |
| ) | |
| response_text = extract_response_text(completion) | |
| except Exception as exc: # noqa: BLE001 | |
| warn_once( | |
| f"model:{task_name}", | |
| f"[{task_name}] Model request failed at step {step}: {exc}. Using heuristic actions.", | |
| ) | |
| return safe_heuristic_action(observation, f"model request failure on {task_name} step {step}") | |
| if DEBUG: | |
| print(f"[{task_name}] model response: {response_text[:300]}") | |
| try: | |
| return parse_model_action(response_text, observation) | |
| except Exception as exc: # noqa: BLE001 | |
| warn_once( | |
| f"parse:{task_name}", | |
| f"[{task_name}] Failed to parse model response at step {step}: {exc}. " | |
| "Using heuristic actions.", | |
| ) | |
| return safe_heuristic_action(observation, f"parse failure on {task_name} step {step}") | |
| async def run_task(task_name: str, client: Optional[OpenAI]) -> TaskRunResult: | |
| env = SatelliteTaskEnv(task_name=task_name) | |
| task = TASK_TYPES[task_name]() | |
| grader = TaskGrader(task) | |
| history: List[str] = [] | |
| observation = env.reset() | |
| state = env.state() | |
| step_limit = state.max_steps | |
| print(f"[START] task={task_name}", flush=True) | |
| for step in range(1, step_limit + 1): | |
| obs_dict = observation_to_dict(observation) | |
| actions = choose_actions( | |
| client, | |
| task_name, | |
| step, | |
| obs_dict, | |
| history, | |
| observation.total_reward, | |
| ) | |
| observation, reward, done, info = env.step( | |
| SatelliteAction(satellite_actions=actions) | |
| ) | |
| reward_value = float(reward.value) | |
| history.append(f"step {step}: {actions} -> reward {reward_value:+.2f}") | |
| print(f"[STEP] step={step} reward={reward_value:.4f} total={observation.total_reward:.4f}",flush=True) | |
| if REQUEST_DELAY > 0 and step < step_limit and not done: | |
| time.sleep(REQUEST_DELAY) | |
| if done: | |
| break | |
| else: | |
| pass | |
| final_state = env.state() | |
| metrics = {key: float(value) for key, value in final_state.metrics.items()} | |
| score = grader.grade_episode(env) | |
| # epsilon = 1e-6 | |
| # score = max(epsilon, min(1.0 - epsilon, score)) | |
| print(f"[END] task={task_name} score={score:.6f} steps={final_state.step_count} done={final_state.done}",flush=True) | |
| return TaskRunResult( | |
| task_name=task_name, | |
| score=score, | |
| total_reward=final_state.total_reward, | |
| steps=final_state.step_count, | |
| done=final_state.done, | |
| metrics=metrics, | |
| ) | |
| def print_summary(results: List[TaskRunResult]) -> None: | |
| aggregate = sum(result.score for result in results) / len(results) | |
| print("\nInference Summary") | |
| print("=" * 60) | |
| for result in results: | |
| print( | |
| f"{result.task_name:<6} score={result.score:.4f} " | |
| f"reward={result.total_reward:.2f} steps={result.steps} done={result.done}" | |
| ) | |
| print("-" * 60) | |
| print(f"aggregate_score={aggregate:.4f}") | |
| print("=" * 60) | |
| async def async_main() -> None: | |
| client = build_client() | |
| results = [] | |
| for task_name in TASK_ORDER: | |
| results.append(await run_task(task_name, client)) | |
| # print_summary(results) | |
| def main() -> None: | |
| try: | |
| asyncio.run(async_main()) | |
| except KeyboardInterrupt: | |
| print("\nInference interrupted.", file=sys.stderr) | |
| raise SystemExit(130) from None | |
| except Exception as exc: # noqa: BLE001 | |
| print(f"\nInference failed: {exc}", file=sys.stderr) | |
| raise SystemExit(1) from None | |
| if __name__ == "__main__": | |
| main() | |