Skip to content
Open
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
1 change: 1 addition & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand Down
30 changes: 23 additions & 7 deletions rio/_tests/test_mw.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
import multiprocessing as mp
import queue
import random
import socket
import unittest
from enum import Enum, auto

Expand Down Expand Up @@ -133,22 +134,36 @@ 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):
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

Expand All @@ -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:
Expand Down
16 changes: 15 additions & 1 deletion rio/_tests/test_policy.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
27 changes: 26 additions & 1 deletion rio/_tests/test_policy_node.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
# SPDX-License-Identifier: Apache-2.0

import os
import socket
import time
from dataclasses import dataclass

Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -84,23 +86,46 @@ 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 = {
"policy": policy,
"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)
Expand Down
14 changes: 11 additions & 3 deletions rio/_tests/test_safety.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
"""

import multiprocessing as mp
import socket
import time

import numpy as np
Expand All @@ -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",
Expand All @@ -38,6 +45,7 @@ def _kwargs():
"freq": 50,
"max_buffer_size": 30,
"camera_keys": ["camera_1"],
"shm_addr": shm_addr,
}


Expand All @@ -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)

Expand Down Expand Up @@ -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)

Expand Down
18 changes: 17 additions & 1 deletion rio/envs/factory.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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"(?<!^)(?=[A-Z])", "_", class_name).lower()


def make_policy(policy_name, policy_kwargs):
module = import_module(f"rio.policies.{policy_name.lower()}")
module_names = [policy_name.lower(), _module_name_from_class_name(policy_name)]
for module_name in dict.fromkeys(module_names):
try:
module = import_module(f"rio.policies.{module_name}")
break
except ModuleNotFoundError as exc:
if exc.name != f"rio.policies.{module_name}":
raise
else:
raise ImportError(policy_name)

PolicyClass = getattr(module, policy_name)
if PolicyClass is None:
raise ImportError(policy_name)
Expand Down
9 changes: 5 additions & 4 deletions rio/policies/policy_interface.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@ def __init__(
max_buffer_size: int = 30,
chunk_request_threshold: float = 0.75,
camera_keys: list[str] | None = None,
**kwargs,
):
self.policy = policy
self.instruction = instruction
Expand All @@ -47,16 +48,16 @@ def __init__(
self.chunk_request_threshold = chunk_request_threshold
self.camera_keys = camera_keys

super().__init__()
super().__init__(freq=freq, max_buffer_size=max_buffer_size, **kwargs)

# NOTE: Defines request/pub schemas in post_init following other constructors
def __post_init__(self):
if len(self.camera_keys) != len(self.resolutions):
logger.error("Length of camera_keys must match length of resolutions")
raise ValueError("Length of camera_keys must match length of resolutions")
if self.camera_keys is None:
logger.warning("camera_keys is None, defaulting to camera_1, camera_2, ...")
self.camera_keys = [f"camera_{i + 1}" for i in range(len(self.resolutions))]
if len(self.camera_keys) != len(self.resolutions):
logger.error("Length of camera_keys must match length of resolutions")
raise ValueError("Length of camera_keys must match length of resolutions")

obs_schema = {"proprio": np.zeros(shape=(self.proprio_dim,), dtype=np.float32)}
for i, resolution in enumerate(self.resolutions):
Expand Down