diff --git a/pyproject.toml b/pyproject.toml index 79213f4..25fcd37 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -36,6 +36,7 @@ rio-list-cameras = "rio._scripts.list_available:list_cameras" rio-list-interfaces = "rio._scripts.list_available:list_interfaces" [project.optional-dependencies] +dev = ["pytest", "pytest-xdist"] visualizers = ["rerun-sdk", "mujoco", "robot_descriptions", "rerun-loader-mjcf"] [dependency-groups] diff --git a/rio/_tests/test_mw.py b/rio/_tests/test_mw.py index c22f3a0..b5d5284 100644 --- a/rio/_tests/test_mw.py +++ b/rio/_tests/test_mw.py @@ -4,6 +4,7 @@ import multiprocessing as mp import queue import random +import socket import unittest from enum import Enum, auto @@ -133,12 +134,29 @@ def payload_client_factory(middleware, **kwargs): return ClientFactory(middleware, PayloadTest, **kwargs) +def middleware_kwargs(middleware, addr, verbose, timeout, freq): + kwargs = {"verbose": verbose, "timeout": timeout, "freq": freq} + if middleware == "Shm": + kwargs["shm_addr"] = addr + elif middleware != "Thread": + kwargs["addr"] = addr + return kwargs + + def create_server_fn(middleware, addr, verbose, timeout, freq): - return lambda: payload_server_factory(middleware, addr=addr, verbose=verbose, timeout=timeout, freq=freq) + kwargs = middleware_kwargs(middleware, addr, verbose, timeout, freq) + return lambda: payload_server_factory(middleware, **kwargs) def create_client_fn(middleware, addr, verbose, timeout, freq): - return lambda: payload_client_factory(middleware, addr=addr, verbose=verbose, timeout=timeout, freq=freq) + kwargs = middleware_kwargs(middleware, addr, verbose, timeout, freq) + return lambda: payload_client_factory(middleware, **kwargs) + + +def unused_addr(): + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock: + sock.bind(("127.0.0.1", 0)) + return f"127.0.0.1:{sock.getsockname()[1]}" class TestMiddleware(unittest.TestCase): @@ -146,9 +164,6 @@ def setUp(self): self.middlewares = ["Zenoh", "Shm", "Thread", "Portal", "ZeroRpc"] self.freq = 50 self.seed = 42 - self._host = "127.0.0.1" - self._port = 7447 - self.addr = f"{self._host}:{self._port}" self.verbose = False self.timeout = 1.0 @@ -159,8 +174,9 @@ def setUp(self): def _test_simple_msg(self, middleware): """Send simple messages through the middleware""" - server_fn = create_server_fn(middleware, self.addr, self.verbose, self.timeout, self.freq) - client_fn = create_client_fn(middleware, self.addr, self.verbose, self.timeout, self.freq) + addr = unused_addr() + server_fn = create_server_fn(middleware, addr, self.verbose, self.timeout, self.freq) + client_fn = create_client_fn(middleware, addr, self.verbose, self.timeout, self.freq) with ServerManager(middleware, [server_fn]): with client_fn() as payload: diff --git a/rio/_tests/test_policy.py b/rio/_tests/test_policy.py index e471d5e..9276a69 100644 --- a/rio/_tests/test_policy.py +++ b/rio/_tests/test_policy.py @@ -58,11 +58,25 @@ def make_dummy_observation(test_cfg): return obs +def make_policy_or_skip(policy_name, policy_kwargs): + try: + policy = make_policy(policy_name, policy_kwargs) + except ImportError as exc: + pytest.skip(f"{policy_name} dependencies are not installed: {exc}") + + required_methods = ("construct_policy", "set_instruction", "get_action") + missing = [name for name in required_methods if not callable(getattr(policy, name, None))] + if missing: + pytest.skip(f"{policy_name} does not implement the policy inference API: missing {missing}") + + return policy + + @pytest.mark.gpu @pytest.mark.parametrize("policy_name", POLICIES_TO_TEST) def test_simple_inference(policy_name, policy_cfg, make_dummy_observation): # Build policy - policy = make_policy(policy_name, vars(policy_cfg)) + policy = make_policy_or_skip(policy_name, vars(policy_cfg)) policy.construct_policy() policy.set_instruction("Move the robot arm to the left.") # Create dummy observation diff --git a/rio/_tests/test_policy_node.py b/rio/_tests/test_policy_node.py index 4c6d748..66a4f03 100644 --- a/rio/_tests/test_policy_node.py +++ b/rio/_tests/test_policy_node.py @@ -2,6 +2,7 @@ # SPDX-License-Identifier: Apache-2.0 import os +import socket import time from dataclasses import dataclass @@ -36,6 +37,7 @@ class PolicyInterfaceConfig: instruction: str = "Move the robot arm to the left." resolutions: list[tuple[int, int]] = None action_dim: int = 6 + proprio_dim: int = 6 chunk_size: int = 50 use_rtc: bool = False freq: int = 100 @@ -84,12 +86,32 @@ def make_dummy_observation(test_cfg): return obs +def shm_addr(): + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock: + sock.bind(("127.0.0.1", 0)) + return f"127.0.0.1:{sock.getsockname()[1]}" + + +def make_policy_or_skip(policy_name, policy_kwargs): + try: + policy = make_policy(policy_name, policy_kwargs) + except ImportError as exc: + pytest.skip(f"{policy_name} dependencies are not installed: {exc}") + + required_methods = ("construct_policy", "set_instruction", "inference") + missing = [name for name in required_methods if not callable(getattr(policy, name, None))] + if missing: + pytest.skip(f"{policy_name} does not implement the policy interface API: missing {missing}") + + return policy + + @pytest.mark.gpu @pytest.mark.integration @pytest.mark.parametrize("policy_name", POLICIES_TO_TEST) def test_policy_interface(policy_name, policy_cfg, policy_interface_cfg, test_cfg, make_dummy_observation): # 1. Instantiate policy wrapper - policy = make_policy(policy_name, vars(policy_cfg)) + policy = make_policy_or_skip(policy_name, vars(policy_cfg)) # 2. Instantiate policy node (server and client factories) policy_interface_kwargs = { @@ -97,10 +119,13 @@ def test_policy_interface(policy_name, policy_cfg, policy_interface_cfg, test_cf "instruction": policy_interface_cfg.instruction, "resolutions": policy_interface_cfg.resolutions, "action_dim": policy_interface_cfg.action_dim, + "proprio_dim": policy_interface_cfg.proprio_dim, "chunk_size": policy_interface_cfg.chunk_size, "use_rtc": policy_interface_cfg.use_rtc, "freq": policy_interface_cfg.freq, "max_buffer_size": policy_interface_cfg.max_buffer_size, + "camera_keys": [f"camera{i + 1}" for i in range(test_cfg.num_cams)], + "shm_addr": shm_addr(), } server = lambda: PolicyInterfaceServer(test_cfg.mw, **policy_interface_kwargs) diff --git a/rio/_tests/test_safety.py b/rio/_tests/test_safety.py index 5b336b6..f5afa96 100644 --- a/rio/_tests/test_safety.py +++ b/rio/_tests/test_safety.py @@ -9,6 +9,7 @@ """ import multiprocessing as mp +import socket import time import numpy as np @@ -27,7 +28,13 @@ CHUNK = 8 -def _kwargs(): +def _shm_addr(): + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock: + sock.bind(("127.0.0.1", 0)) + return f"127.0.0.1:{sock.getsockname()[1]}" + + +def _kwargs(shm_addr): return { "policy": Dummy(action_dim=ACTION_DIM, chunk_size=CHUNK), "instruction": "test", @@ -38,6 +45,7 @@ def _kwargs(): "freq": 50, "max_buffer_size": 30, "camera_keys": ["camera_1"], + "shm_addr": shm_addr, } @@ -49,7 +57,7 @@ def _obs(): def test_gap_returns_safe_default_then_recovers(): - kw = _kwargs() + kw = _kwargs(_shm_addr()) server = lambda: PolicyInterfaceServer(MW, **kw) client = lambda: PolicyInterfaceClient(MW, **kw) @@ -77,7 +85,7 @@ def test_gap_returns_safe_default_then_recovers(): def test_get_action_chunk_never_blocks(): - kw = _kwargs() + kw = _kwargs(_shm_addr()) server = lambda: PolicyInterfaceServer(MW, **kw) client = lambda: PolicyInterfaceClient(MW, **kw) diff --git a/rio/envs/factory.py b/rio/envs/factory.py index 9a5f09c..8d05d7b 100644 --- a/rio/envs/factory.py +++ b/rio/envs/factory.py @@ -1,6 +1,7 @@ # SPDX-FileCopyrightText: 2026 RIO Developers # SPDX-License-Identifier: Apache-2.0 +import re from contextlib import ExitStack, contextmanager from dataclasses import asdict, is_dataclass from dataclasses import fields as dataclass_fields @@ -51,8 +52,23 @@ def dataclass_to_dict(dc): return result +def _module_name_from_class_name(class_name: str) -> str: + """Convert a PascalCase class name to its snake_case module name.""" + return re.sub(r"(?