Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
58 changes: 58 additions & 0 deletions aas-gym/src/aas_gym/aas_config.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,58 @@
"""Pure AASEnv configuration validators (no Docker/ROS imports)."""

from __future__ import annotations
import math

ALLOWED_AUTOPILOTS = frozenset({"px4", "ardupilot"})
ALLOWED_ODOM = frozenset({"none", "openvins", "fastlio", "superodom", "mimosa"})
ALLOWED_RENDER_MODES = frozenset({None, "human", "ansi"})


def validate_aas_config(
*,
instance: int = 0,
gym_freq_hz: int = 50,
autopilot: str = "px4",
odom: str = "none",
num_quads: int = 1,
render_mode=None,
) -> dict:
"""Validate constructor knobs; raise ValueError on bad input.

Returns a normalized dict of the checked fields.
"""
if isinstance(instance, bool) or not isinstance(instance, int) or instance < 0:
raise ValueError(f"instance must be a non-negative int, got {instance!r}")
if isinstance(gym_freq_hz, bool) or not isinstance(gym_freq_hz, int) or gym_freq_hz < 1:
raise ValueError(f"gym_freq_hz must be an int >= 1, got {gym_freq_hz!r}")
if autopilot not in ALLOWED_AUTOPILOTS:
raise ValueError(f"autopilot must be one of {sorted(ALLOWED_AUTOPILOTS)}, got {autopilot!r}")
if odom not in ALLOWED_ODOM:
raise ValueError(f"odom must be one of {sorted(ALLOWED_ODOM)}, got {odom!r}")
if isinstance(num_quads, bool) or not isinstance(num_quads, int) or num_quads < 1:
raise ValueError(f"num_quads must be an int >= 1, got {num_quads!r}")
if render_mode not in ALLOWED_RENDER_MODES:
raise ValueError(f"render_mode must be one of human/ansi/None, got {render_mode!r}")
return {
"instance": instance,
"gym_freq_hz": gym_freq_hz,
"autopilot": autopilot,
"odom": odom,
"num_quads": num_quads,
"render_mode": render_mode,
}


def validate_zmq_transport(transport: str) -> str:
if transport not in ("tcp", "ipc"):
raise ValueError(f"ZMQ_TRANSPORT must be 'tcp' or 'ipc', got {transport!r}")
return transport


def finite_episode_seconds(value) -> float:
if isinstance(value, bool) or not isinstance(value, (int, float)):
raise ValueError(f"seconds must be numeric, got {type(value).__name__}")
f = float(value)
if not math.isfinite(f) or f <= 0:
raise ValueError(f"seconds must be finite and > 0, got {f!r}")
return f
14 changes: 14 additions & 0 deletions aas-gym/src/aas_gym/aas_env.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,11 @@

from docker.types import DeviceRequest

try:
from aas_gym.aas_config import validate_aas_config
except ImportError:
from aas_config import validate_aas_config


class AASEnv(gym.Env):
metadata = {"render_modes": ["human", "ansi"]}
Expand All @@ -28,6 +33,15 @@ def __init__(self,
):
super().__init__()

validate_aas_config(
instance=instance,
gym_freq_hz=gym_freq_hz,
autopilot=autopilot,
odom=odom,
num_quads=num_quads,
render_mode=render_mode,
)

self.GYM_FREQ_HZ = gym_freq_hz
self.GYM_INIT_DURATION = 80.0 # Seconds to run unpaused during reset (seconds)
self.MAX_EPISODE_LENGTH_SEC = 300.0 # Max episode length in seconds (excluding init duration)
Expand Down
53 changes: 53 additions & 0 deletions aas-gym/src/aas_gym/test_aas_config.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,53 @@
#!/usr/bin/env python3
import math
import unittest

try:
from aas_gym.aas_config import (
validate_aas_config,
validate_zmq_transport,
finite_episode_seconds,
)
except ImportError:
from aas_config import (
validate_aas_config,
validate_zmq_transport,
finite_episode_seconds,
)


class TestAasConfig(unittest.TestCase):
def test_defaults_ok(self):
cfg = validate_aas_config()
self.assertEqual(cfg["autopilot"], "px4")
self.assertEqual(cfg["num_quads"], 1)

def test_bad_num_quads(self):
with self.assertRaises(ValueError):
validate_aas_config(num_quads=0)
with self.assertRaises(ValueError):
validate_aas_config(num_quads=-1)

def test_bad_odom(self):
with self.assertRaises(ValueError):
validate_aas_config(odom="vins-fusion")

def test_bad_freq(self):
with self.assertRaises(ValueError):
validate_aas_config(gym_freq_hz=0)

def test_zmq(self):
self.assertEqual(validate_zmq_transport("tcp"), "tcp")
with self.assertRaises(ValueError):
validate_zmq_transport("udp")

def test_episode_seconds(self):
self.assertEqual(finite_episode_seconds(300), 300.0)
for bad in (0, -1, math.nan, math.inf):
with self.subTest(bad=bad):
with self.assertRaises(ValueError):
finite_episode_seconds(bad)


if __name__ == "__main__":
unittest.main()
Loading