diff --git a/dimos/core/coordination/blueprint_config/parser.py b/dimos/core/coordination/blueprint_config/parser.py index ce8338ec66..a97422040f 100644 --- a/dimos/core/coordination/blueprint_config/parser.py +++ b/dimos/core/coordination/blueprint_config/parser.py @@ -75,6 +75,7 @@ plain, plain_mapping, snapshot_mapping, + validated_model_values, ) from dimos.core.coordination.blueprints import ( Blueprint, @@ -421,7 +422,7 @@ def _validate_modules( raise BlueprintConfigError( format_validation_error(module.atom.name, error) ) from error - dumped = model.model_dump(mode="python", exclude_unset=True) + dumped = validated_model_values(model) dumped.pop("g", None) dumped.pop("instance_name", None) parsed[module.atom.name] = dumped diff --git a/dimos/core/coordination/blueprint_config/test_parser.py b/dimos/core/coordination/blueprint_config/test_parser.py index 6b24322489..bcfdf272fb 100644 --- a/dimos/core/coordination/blueprint_config/test_parser.py +++ b/dimos/core/coordination/blueprint_config/test_parser.py @@ -13,7 +13,9 @@ # limitations under the License. from collections.abc import Callable +from dataclasses import dataclass from pathlib import Path +import pickle from typing import Annotated, Any, Literal from pydantic import BaseModel, Field @@ -452,6 +454,14 @@ def __str__(self) -> str: return "Anchor:\n multi\n line" +@dataclass(frozen=True) +class CallableAnchor: + prefix: str + + def __call__(self, value: Any) -> str: + return f"{self.prefix}:{value}" + + class ArbitraryConfig(ModuleConfig): scaling: Anchor = Field(default_factory=Anchor) hybrid: Anchor | str = "fallback" @@ -490,6 +500,18 @@ def test_blueprint_pinned_arbitrary_value_survives_filtering() -> None: assert isinstance(parsed.module_kwargs("arbitrarymodule")["scaling"], Anchor) +def test_blueprint_pinned_callable_dataclass_survives_worker_serialization() -> None: + blueprint = ArbitraryModule.blueprint(handlers={"scene": CallableAnchor("render")}) + + parsed = BlueprintConfigParser(blueprint).parse(environ={}) + worker_kwargs = pickle.loads(pickle.dumps(parsed.module_kwargs("arbitrarymodule"))) + worker_config = ArbitraryConfig.model_validate(worker_kwargs) + handler = worker_config.handlers["scene"] + + assert isinstance(handler, CallableAnchor) + assert handler("apartment") == "render:apartment" + + def test_format_help_uses_nested_parent_default_instance() -> None: class NestedRequiredConfig(BaseModel): value: int diff --git a/dimos/core/coordination/blueprint_config/values.py b/dimos/core/coordination/blueprint_config/values.py index 39cab50442..b5654d6bbc 100644 --- a/dimos/core/coordination/blueprint_config/values.py +++ b/dimos/core/coordination/blueprint_config/values.py @@ -55,6 +55,31 @@ def plain(value: Any) -> Any: return _copy_opaque(value) +def validated_model_values(model: BaseModel) -> dict[str, Any]: + """Copy explicitly set validated fields without serializing runtime objects.""" + return { + name: _validated_value(getattr(model, name)) + for name in type(model).model_fields + if name in model.model_fields_set + } + + +def _validated_value(value: Any) -> Any: + if isinstance(value, BaseModel): + return validated_model_values(value) + if isinstance(value, Mapping): + return {_copy_opaque(key): _validated_value(item) for key, item in value.items()} + if isinstance(value, list): + return [_validated_value(item) for item in value] + if isinstance(value, tuple): + return tuple(_validated_value(item) for item in value) + if isinstance(value, set): + return {_validated_value(item) for item in value} + if isinstance(value, frozenset): + return frozenset(_validated_value(item) for item in value) + return _copy_opaque(value) + + def deep_merge(destination: dict[str, Any], incoming: Mapping[str, Any]) -> None: for key, value in incoming.items(): if key in destination and isinstance(destination[key], dict) and isinstance(value, Mapping): diff --git a/dimos/core/coordination/coordinator_rpc.py b/dimos/core/coordination/coordinator_rpc.py index ac4fe9960b..3c64fdc85c 100644 --- a/dimos/core/coordination/coordinator_rpc.py +++ b/dimos/core/coordination/coordinator_rpc.py @@ -19,6 +19,8 @@ from dimos.core.global_config import global_config from dimos.core.transport_factory import rpc_backend +from dimos.protocol.rpc.zenohrpc import ZenohRPC +from dimos.protocol.service.zenohservice import ZENOH_LOCAL_ROUTER_ENDPOINT from dimos.utils.logging_config import setup_logger if TYPE_CHECKING: @@ -51,7 +53,12 @@ def serve(cls, coordinator: RPCInspectable) -> CoordinatorRPC: @classmethod def connect(cls, *, timeout: float) -> CoordinatorRPC: """Attach to a running Coordinator, raising `TimeoutError` if none answers.""" - rpc = rpc_backend()() + backend = rpc_backend() + rpc = ( + ZenohRPC(mode="client", connect=[ZENOH_LOCAL_ROUTER_ENDPOINT]) + if backend is ZenohRPC + else backend() + ) rpc.start() client = cls(rpc) deadline = time.monotonic() + timeout diff --git a/dimos/core/coordination/module_coordinator.py b/dimos/core/coordination/module_coordinator.py index 24b7a70ce2..6187d73a99 100644 --- a/dimos/core/coordination/module_coordinator.py +++ b/dimos/core/coordination/module_coordinator.py @@ -19,6 +19,7 @@ import dataclasses import importlib import inspect +import os import shutil import sys import threading @@ -41,6 +42,11 @@ pZenohTransport, ) from dimos.core.transport_factory import make_transport +from dimos.protocol.service.zenohservice import ( + ZENOH_LOCAL_ROUTER_ENDPOINT, + ZENOH_ROUTER_ENDPOINT_ENV, + ZenohRouter, +) from dimos.spec.utils import is_spec, spec_annotation_compliance, spec_structural_compliance from dimos.utils.generic import short_id from dimos.utils.logging_config import setup_logger @@ -89,13 +95,24 @@ def __init__( self._modules_lock = threading.RLock() self._rpc_lock = threading.RLock() self._coordinator_rpc: CoordinatorRPC | None = None + self._zenoh_router: ZenohRouter | None = None + self._previous_zenoh_router_endpoint: str | None = None def start(self) -> None: from dimos.core.o3dpickle import register_picklers register_picklers() - for m in self._managers.values(): - m.start() + if self._global_config.transport == "zenoh": + self._zenoh_router = ZenohRouter() + self._zenoh_router.start() + self._previous_zenoh_router_endpoint = os.environ.get(ZENOH_ROUTER_ENDPOINT_ENV) + os.environ[ZENOH_ROUTER_ENDPOINT_ENV] = ZENOH_LOCAL_ROUTER_ENDPOINT + try: + for m in self._managers.values(): + m.start() + except BaseException: + self._stop_zenoh_router() + raise self._started = True def stop(self) -> None: @@ -119,9 +136,21 @@ def _stop_manager(m: WorkerManager) -> None: logger.error("Error stopping manager", manager=type(m).__name__, exc_info=True) safe_thread_map(tuple(self._managers.values()), _stop_manager) + self._stop_zenoh_router() + + def _stop_zenoh_router(self) -> None: + if self._zenoh_router is not None: + self._zenoh_router.stop() + self._zenoh_router = None + if self._global_config.transport != "zenoh": + return + if self._previous_zenoh_router_endpoint is None: + os.environ.pop(ZENOH_ROUTER_ENDPOINT_ENV, None) + else: + os.environ[ZENOH_ROUTER_ENDPOINT_ENV] = self._previous_zenoh_router_endpoint def start_rpc_service(self) -> None: - """Expose the coordinator's API as @rpc methods over LCM.""" + """Expose the coordinator's API over the configured RPC transport.""" with self._rpc_lock: if self._coordinator_rpc is not None: return diff --git a/dimos/core/global_config.py b/dimos/core/global_config.py index f5cce34f11..1f1865c19d 100644 --- a/dimos/core/global_config.py +++ b/dimos/core/global_config.py @@ -53,6 +53,7 @@ class GlobalConfig(BaseSettings): can_port: str | None = None device_path: str | None = None # device path for real robot (e.g. /dev/ttyUSB0) simulation: str = "" + simulation_provider: str = "" replay: bool = False replay_db: str = "go2_short" new_memory: bool = False diff --git a/dimos/core/native_module.py b/dimos/core/native_module.py index 461bd4d0ed..d7b712a66f 100644 --- a/dimos/core/native_module.py +++ b/dimos/core/native_module.py @@ -60,6 +60,10 @@ class MyCppModule(NativeModule): from dimos.core.core import rpc from dimos.core.global_config import global_config from dimos.core.module import Module, ModuleConfig +from dimos.protocol.service.zenohservice import ( + ZENOH_LOCAL_ROUTER_ENDPOINT, + ZENOH_ROUTER_ENDPOINT_ENV, +) from dimos.utils.logging_config import setup_logger if sys.platform.startswith("linux"): @@ -254,6 +258,8 @@ def start(self) -> None: # set transport so native modules know which one to spawn env["DIMOS_TRANSPORT"] = global_config.transport + if global_config.transport == "zenoh": + env[ZENOH_ROUTER_ENDPOINT_ENV] = ZENOH_LOCAL_ROUTER_ENDPOINT # set Rust logging to match Python level env["RUST_LOG"] = _PYTHON_TO_RUST_LEVELS.get( diff --git a/dimos/hardware/manipulators/sim/adapter.py b/dimos/hardware/manipulators/sim/adapter.py index 6b00bbf52e..d6e9977a88 100644 --- a/dimos/hardware/manipulators/sim/adapter.py +++ b/dimos/hardware/manipulators/sim/adapter.py @@ -27,7 +27,7 @@ JointLimits, ManipulatorInfo, ) -from dimos.simulation.engines.mujoco_shm import ( +from dimos.hardware.simulation.shared_memory import ( ManipShmReader, shm_key_from_path, ) diff --git a/dimos/hardware/manipulators/sim/test_shm_adapter.py b/dimos/hardware/manipulators/sim/test_shm_adapter.py index 5c1c0d17e2..f1e23504c6 100644 --- a/dimos/hardware/manipulators/sim/test_shm_adapter.py +++ b/dimos/hardware/manipulators/sim/test_shm_adapter.py @@ -23,7 +23,7 @@ import dimos.hardware.manipulators.sim.adapter as adapter_mod from dimos.hardware.manipulators.sim.adapter import ShmMujocoAdapter from dimos.hardware.manipulators.spec import ControlMode, ManipulatorAdapter -from dimos.simulation.engines.mujoco_shm import ManipShmWriter +from dimos.hardware.simulation.shared_memory import ManipShmWriter ARM_DOF = 7 diff --git a/dimos/hardware/simulation/shared_memory.py b/dimos/hardware/simulation/shared_memory.py new file mode 100644 index 0000000000..f6a934d94b --- /dev/null +++ b/dimos/hardware/simulation/shared_memory.py @@ -0,0 +1,504 @@ +# Copyright 2025-2026 Dimensional Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Shared-memory buffers for sim-manipulator IPC. + +Layout for exchanging joint state and commands between ``MujocoSimModule`` +(which owns the physics engine) and ``ShmMujocoAdapter`` (which plugs into +ControlCoordinator). Modeled after ``dimos.simulation.mujoco.shared_memory`` +(the Go2 SHM pattern). + +Names are deterministic: both sides derive them from the resolved MJCF path, +so no name exchange over RPC is needed. The sim module creates the buffers +and signals ``ready``; the adapter attaches to them by name. +""" + +from __future__ import annotations + +from dataclasses import dataclass +import hashlib +from multiprocessing.shared_memory import SharedMemory +from pathlib import Path +from typing import Any + +import numpy as np +from numpy.typing import NDArray + +from dimos.utils.logging_config import setup_logger + +logger = setup_logger() + +# Upper bound on joint count per sim. Manipulators use <=10; humanoids +# (Unitree G1: 29) push higher. 32 leaves headroom while keeping all +# per-joint buffers tiny (32 floats = 256 B). +MAX_JOINTS = 32 +_FLOAT_BYTES = 8 # float64 +_INT32_BYTES = 4 + +# IMU layout: quat (4) + gyro (3) + accel (3) = 10 floats. +_IMU_FLOATS = 10 + +_joint_array_size = MAX_JOINTS * _FLOAT_BYTES # float64 array + +# Element counts for control and sequence arrays. +_NUM_CTRL_FIELDS = 4 # [ready, stop, command_mode, num_joints] +_NUM_SEQ_COUNTERS = 12 # one per buffer type (manipulator + WB additions) + +# Buffer sizes (in bytes). +# Keys are short to stay under macOS PSHMNAMLEN (31 bytes). +_shm_sizes = { + # Manipulator-shared layout + "pos": _joint_array_size, + "vel": _joint_array_size, + "eff": _joint_array_size, + "pos_t": _joint_array_size, + "vel_t": _joint_array_size, + "grp": 2 * _FLOAT_BYTES, # [gripper_position, gripper_target] + # Whole-body additions (unused by manipulator path). + "imu": _IMU_FLOATS * _FLOAT_BYTES, # [w,x,y,z, gx,gy,gz, ax,ay,az] + "kp_t": _joint_array_size, # per-joint position-gain target + "kd_t": _joint_array_size, # per-joint velocity-gain target + "tau_t": _joint_array_size, # per-joint feedforward torque + # Bookkeeping + "seq": _NUM_SEQ_COUNTERS * _FLOAT_BYTES, # int64 counters + "ctl": _NUM_CTRL_FIELDS * _INT32_BYTES, # [ready, stop, command_mode, num_joints] +} + +# Sequence counter indices. +SEQ_POSITIONS = 0 +SEQ_VELOCITIES = 1 +SEQ_EFFORTS = 2 +SEQ_POSITION_CMD = 3 +SEQ_VELOCITY_CMD = 4 +SEQ_GRIPPER_STATE = 5 +SEQ_GRIPPER_CMD = 6 +# Whole-body additions +SEQ_IMU = 7 +SEQ_KP_CMD = 8 +SEQ_KD_CMD = 9 +SEQ_TAU_CMD = 10 + +# Control indices. +CTRL_READY = 0 +CTRL_STOP = 1 +CTRL_COMMAND_MODE = 2 +CTRL_NUM_JOINTS = 3 + +# Command modes. +CMD_MODE_POSITION = 0 +CMD_MODE_VELOCITY = 1 +# Whole-body PD-with-feedforward: ctrl = kp*(q_t - q) + kd*(0 - dq) + tau_t. +# Per-step kp/kd lets a policy retune gains online if it wants to. +CMD_MODE_PD_TAU = 2 + +_NAME_PREFIX = "dmjm" + + +def shm_key_from_path(config_path: Path | str) -> str: + """Derive a deterministic short key from a simulation model path. + + Both simulation provider and adapter compute the same key from the same path, + so SHM buffer names can be agreed upon without an RPC round-trip. + """ + resolved = str(Path(config_path).expanduser().resolve()) + return hashlib.md5(resolved.encode("utf-8")).hexdigest()[:12] + + +def _buffer_name(key: str, buffer: str) -> str: + return f"{_NAME_PREFIX}_{key}_{buffer}" + + +@dataclass(frozen=True) +class ManipShmSet: + """Frozen set of named SharedMemory buffers for sim <-> adapter IPC. + + Despite the name (kept for backward compat with existing manipulator + consumers), the layout now also covers whole-body needs: IMU, per-joint + PD gain commands, and per-joint feedforward torque commands. The + extra buffers are unused by the manipulator path. + """ + + pos: SharedMemory + vel: SharedMemory + eff: SharedMemory + pos_t: SharedMemory + vel_t: SharedMemory + grp: SharedMemory + # Whole-body additions + imu: SharedMemory + kp_t: SharedMemory + kd_t: SharedMemory + tau_t: SharedMemory + # Bookkeeping + seq: SharedMemory + ctl: SharedMemory + + @classmethod + def create(cls, key: str) -> ManipShmSet: + """Create new SHM buffers with deterministic names derived from *key*""" + buffers: dict[str, SharedMemory] = {} + for buffer_name, size in _shm_sizes.items(): + name = _buffer_name(key, buffer_name) + try: + stale = SharedMemory(name=name) + stale.close() + try: + stale.unlink() + logger.info("ManipShmSet: unlinked stale SHM", name=name) + except FileNotFoundError: + pass + except FileNotFoundError: + pass + buffers[buffer_name] = SharedMemory(create=True, size=size, name=name) + return cls(**buffers) + + @classmethod + def attach(cls, key: str) -> ManipShmSet: + """Attach to existing SHM buffers created by the sim side.""" + buffers: dict[str, SharedMemory] = {} + for buffer_name in _shm_sizes: + name = _buffer_name(key, buffer_name) + buffers[buffer_name] = SharedMemory(name=name) + return cls(**buffers) + + def as_list(self) -> list[SharedMemory]: + return [getattr(self, k) for k in _shm_sizes] + + +class ManipShmWriter: + """Sim-side handle: writes joint state, reads command targets. + Owned by the active simulation provider. Creates the SHM buffers on init and + unlinks them on cleanup. + """ + + shm: ManipShmSet + + def __init__(self, key: str) -> None: + self.shm = ManipShmSet.create(key) + self._last_pos_cmd_seq = 0 + self._last_vel_cmd_seq = 0 + self._last_gripper_cmd_seq = 0 + self._last_kp_cmd_seq = 0 + self._last_kd_cmd_seq = 0 + self._last_tau_cmd_seq = 0 + # Zero everything. + for buf in self.shm.as_list(): + np.ndarray((buf.size,), dtype=np.uint8, buffer=buf.buf)[:] = 0 + + def write_joint_state( + self, + positions: list[float], + velocities: list[float], + efforts: list[float], + ) -> None: + n = min(len(positions), MAX_JOINTS) + pos_arr = self._array(self.shm.pos, MAX_JOINTS, np.float64) + vel_arr = self._array(self.shm.vel, MAX_JOINTS, np.float64) + eff_arr = self._array(self.shm.eff, MAX_JOINTS, np.float64) + pos_arr[:n] = positions[:n] + vel_arr[:n] = velocities[:n] + eff_arr[:n] = efforts[:n] + self._increment_seq(SEQ_POSITIONS) + self._increment_seq(SEQ_VELOCITIES) + self._increment_seq(SEQ_EFFORTS) + + def write_gripper_state(self, position: float) -> None: + arr = self._array(self.shm.grp, 2, np.float64) + arr[0] = position + self._increment_seq(SEQ_GRIPPER_STATE) + + def read_position_command(self, num_joints: int) -> NDArray[np.float64] | None: + """Return a copy of position targets if a new command arrived since last call.""" + seq = self._get_seq(SEQ_POSITION_CMD) + if seq <= self._last_pos_cmd_seq: + return None + self._last_pos_cmd_seq = seq + arr = self._array(self.shm.pos_t, MAX_JOINTS, np.float64) + result: NDArray[np.float64] = arr[:num_joints].copy() + return result + + def read_velocity_command(self, num_joints: int) -> NDArray[np.float64] | None: + seq = self._get_seq(SEQ_VELOCITY_CMD) + if seq <= self._last_vel_cmd_seq: + return None + self._last_vel_cmd_seq = seq + arr = self._array(self.shm.vel_t, MAX_JOINTS, np.float64) + result: NDArray[np.float64] = arr[:num_joints].copy() + return result + + def read_gripper_command(self) -> float | None: + seq = self._get_seq(SEQ_GRIPPER_CMD) + if seq <= self._last_gripper_cmd_seq: + return None + self._last_gripper_cmd_seq = seq + arr = self._array(self.shm.grp, 2, np.float64) + return float(arr[1]) + + def read_command_mode(self) -> int: + return int(self._control()[CTRL_COMMAND_MODE]) + + # Whole-body additions + + def write_imu( + self, + quaternion: tuple[float, float, float, float], + gyroscope: tuple[float, float, float], + accelerometer: tuple[float, float, float], + ) -> None: + """Write IMU sample. Quaternion is (w, x, y, z).""" + arr = self._array(self.shm.imu, _IMU_FLOATS, np.float64) + arr[0:4] = quaternion + arr[4:7] = gyroscope + arr[7:10] = accelerometer + self._increment_seq(SEQ_IMU) + + def read_kp_command(self, num_joints: int) -> NDArray[np.float64] | None: + """Per-joint position-gain target if a new command landed since last call.""" + seq = self._get_seq(SEQ_KP_CMD) + if seq <= self._last_kp_cmd_seq: + return None + self._last_kp_cmd_seq = seq + arr = self._array(self.shm.kp_t, MAX_JOINTS, np.float64) + return arr[:num_joints].copy() + + def read_kd_command(self, num_joints: int) -> NDArray[np.float64] | None: + seq = self._get_seq(SEQ_KD_CMD) + if seq <= self._last_kd_cmd_seq: + return None + self._last_kd_cmd_seq = seq + arr = self._array(self.shm.kd_t, MAX_JOINTS, np.float64) + return arr[:num_joints].copy() + + def read_tau_command(self, num_joints: int) -> NDArray[np.float64] | None: + """Per-joint feedforward torque if a new command landed since last call.""" + seq = self._get_seq(SEQ_TAU_CMD) + if seq <= self._last_tau_cmd_seq: + return None + self._last_tau_cmd_seq = seq + arr = self._array(self.shm.tau_t, MAX_JOINTS, np.float64) + return arr[:num_joints].copy() + + def signal_ready(self, num_joints: int) -> None: + ctrl = self._control() + ctrl[CTRL_NUM_JOINTS] = num_joints + ctrl[CTRL_READY] = 1 + + def signal_stop(self) -> None: + self._control()[CTRL_STOP] = 1 + + def should_stop(self) -> bool: + return bool(self._control()[CTRL_STOP] == 1) + + def cleanup(self) -> None: + for shm in self.shm.as_list(): + try: + shm.close() + except FileNotFoundError: + pass # already detached + except OSError as exc: + logger.warning("SHM close failed", name=shm.name, error=str(exc)) + try: + shm.unlink() + except FileNotFoundError: + pass # already unlinked (e.g. cleanup called twice) + except OSError as exc: + logger.warning("SHM unlink failed", name=shm.name, error=str(exc)) + + def _array(self, buf: SharedMemory, n: int, dtype: Any) -> NDArray[Any]: + return np.ndarray((n,), dtype=dtype, buffer=buf.buf) + + def _control(self) -> NDArray[np.int32]: + return np.ndarray((_NUM_CTRL_FIELDS,), dtype=np.int32, buffer=self.shm.ctl.buf) + + def _increment_seq(self, index: int) -> None: + seq_arr = np.ndarray((_NUM_SEQ_COUNTERS,), dtype=np.int64, buffer=self.shm.seq.buf) + seq_arr[index] += 1 + + def _get_seq(self, index: int) -> int: + seq_arr = np.ndarray((_NUM_SEQ_COUNTERS,), dtype=np.int64, buffer=self.shm.seq.buf) + return int(seq_arr[index]) + + +class ManipShmReader: + """Adapter-side handle: reads joint state, writes command targets. + + Owned by ``ShmMujocoAdapter``. Attaches to existing buffers created by + the sim module; does not unlink them on cleanup. + """ + + shm: ManipShmSet + + def __init__(self, key: str) -> None: + self.shm = ManipShmSet.attach(key) + + def read_positions(self, num_joints: int) -> list[float]: + arr = np.ndarray((MAX_JOINTS,), dtype=np.float64, buffer=self.shm.pos.buf) + return [float(x) for x in arr[:num_joints]] + + def read_velocities(self, num_joints: int) -> list[float]: + arr = np.ndarray((MAX_JOINTS,), dtype=np.float64, buffer=self.shm.vel.buf) + return [float(x) for x in arr[:num_joints]] + + def read_efforts(self, num_joints: int) -> list[float]: + arr = np.ndarray((MAX_JOINTS,), dtype=np.float64, buffer=self.shm.eff.buf) + return [float(x) for x in arr[:num_joints]] + + def read_gripper_position(self) -> float: + arr = np.ndarray((2,), dtype=np.float64, buffer=self.shm.grp.buf) + return float(arr[0]) + + def write_position_command(self, positions: list[float]) -> None: + n = min(len(positions), MAX_JOINTS) + arr = np.ndarray((MAX_JOINTS,), dtype=np.float64, buffer=self.shm.pos_t.buf) + arr[:n] = positions[:n] + self._set_command_mode(CMD_MODE_POSITION) + self._increment_seq(SEQ_POSITION_CMD) + + def write_velocity_command(self, velocities: list[float]) -> None: + n = min(len(velocities), MAX_JOINTS) + arr = np.ndarray((MAX_JOINTS,), dtype=np.float64, buffer=self.shm.vel_t.buf) + arr[:n] = velocities[:n] + self._set_command_mode(CMD_MODE_VELOCITY) + self._increment_seq(SEQ_VELOCITY_CMD) + + def write_gripper_command(self, position: float) -> None: + arr = np.ndarray((2,), dtype=np.float64, buffer=self.shm.grp.buf) + arr[1] = position + self._increment_seq(SEQ_GRIPPER_CMD) + + # Whole-body additions + + def read_imu( + self, + ) -> tuple[ + tuple[float, float, float, float], + tuple[float, float, float], + tuple[float, float, float], + ]: + """Read IMU sample: ((qw, qx, qy, qz), (gx, gy, gz), (ax, ay, az)).""" + arr = np.ndarray((_IMU_FLOATS,), dtype=np.float64, buffer=self.shm.imu.buf) + return ( + (float(arr[0]), float(arr[1]), float(arr[2]), float(arr[3])), + (float(arr[4]), float(arr[5]), float(arr[6])), + (float(arr[7]), float(arr[8]), float(arr[9])), + ) + + def write_kp_command(self, kp: list[float]) -> None: + """Per-joint position-gain target. Switches command mode to PD+tau.""" + n = min(len(kp), MAX_JOINTS) + arr = np.ndarray((MAX_JOINTS,), dtype=np.float64, buffer=self.shm.kp_t.buf) + arr[:n] = kp[:n] + self._set_command_mode(CMD_MODE_PD_TAU) + self._increment_seq(SEQ_KP_CMD) + + def write_kd_command(self, kd: list[float]) -> None: + n = min(len(kd), MAX_JOINTS) + arr = np.ndarray((MAX_JOINTS,), dtype=np.float64, buffer=self.shm.kd_t.buf) + arr[:n] = kd[:n] + self._set_command_mode(CMD_MODE_PD_TAU) + self._increment_seq(SEQ_KD_CMD) + + def write_tau_command(self, tau: list[float]) -> None: + """Per-joint feedforward torque, applied on top of PD.""" + n = min(len(tau), MAX_JOINTS) + arr = np.ndarray((MAX_JOINTS,), dtype=np.float64, buffer=self.shm.tau_t.buf) + arr[:n] = tau[:n] + self._set_command_mode(CMD_MODE_PD_TAU) + self._increment_seq(SEQ_TAU_CMD) + + def write_pd_tau_command( + self, + positions: list[float], + kp: list[float], + kd: list[float], + tau: list[float], + ) -> None: + """Write a whole-body PD+tau command without transient mode flips. + + The sim engine runs in a different process, so setting position mode + first and PD mode later creates a small but real race. Write all arrays, + publish PD mode once, then bump the sequence counters. + """ + n_pos = min(len(positions), MAX_JOINTS) + n_kp = min(len(kp), MAX_JOINTS) + n_kd = min(len(kd), MAX_JOINTS) + n_tau = min(len(tau), MAX_JOINTS) + np.ndarray((MAX_JOINTS,), dtype=np.float64, buffer=self.shm.pos_t.buf)[:n_pos] = positions[ + :n_pos + ] + np.ndarray((MAX_JOINTS,), dtype=np.float64, buffer=self.shm.kp_t.buf)[:n_kp] = kp[:n_kp] + np.ndarray((MAX_JOINTS,), dtype=np.float64, buffer=self.shm.kd_t.buf)[:n_kd] = kd[:n_kd] + np.ndarray((MAX_JOINTS,), dtype=np.float64, buffer=self.shm.tau_t.buf)[:n_tau] = tau[:n_tau] + self._set_command_mode(CMD_MODE_PD_TAU) + self._increment_seq(SEQ_KP_CMD) + self._increment_seq(SEQ_KD_CMD) + self._increment_seq(SEQ_TAU_CMD) + # Position is the engine-side trigger for latching a new PD target, + # so publish it last after gains/torque are visible. + self._increment_seq(SEQ_POSITION_CMD) + + def is_ready(self) -> bool: + return bool(self._control()[CTRL_READY] == 1) + + def num_joints(self) -> int: + return int(self._control()[CTRL_NUM_JOINTS]) + + def signal_stop(self) -> None: + self._control()[CTRL_STOP] = 1 + + def cleanup(self) -> None: + for shm in self.shm.as_list(): + try: + shm.close() + except FileNotFoundError: + pass # already detached + except OSError as exc: + logger.warning("SHM close failed", name=shm.name, error=str(exc)) + + def _control(self) -> NDArray[np.int32]: + return np.ndarray((_NUM_CTRL_FIELDS,), dtype=np.int32, buffer=self.shm.ctl.buf) + + def _set_command_mode(self, mode: int) -> None: + self._control()[CTRL_COMMAND_MODE] = mode + + def _increment_seq(self, index: int) -> None: + seq_arr = np.ndarray((_NUM_SEQ_COUNTERS,), dtype=np.int64, buffer=self.shm.seq.buf) + seq_arr[index] += 1 + + +__all__ = [ + "CMD_MODE_PD_TAU", + "CMD_MODE_POSITION", + "CMD_MODE_VELOCITY", + "CTRL_COMMAND_MODE", + "CTRL_NUM_JOINTS", + "CTRL_READY", + "CTRL_STOP", + "MAX_JOINTS", + "SEQ_EFFORTS", + "SEQ_GRIPPER_CMD", + "SEQ_GRIPPER_STATE", + "SEQ_IMU", + "SEQ_KD_CMD", + "SEQ_KP_CMD", + "SEQ_POSITIONS", + "SEQ_POSITION_CMD", + "SEQ_TAU_CMD", + "SEQ_VELOCITIES", + "SEQ_VELOCITY_CMD", + "ManipShmReader", + "ManipShmSet", + "ManipShmWriter", + "shm_key_from_path", +] diff --git a/dimos/protocol/service/test_zenohservice.py b/dimos/protocol/service/test_zenohservice.py index 621c370485..c1406023b4 100644 --- a/dimos/protocol/service/test_zenohservice.py +++ b/dimos/protocol/service/test_zenohservice.py @@ -16,7 +16,13 @@ import pytest -from dimos.protocol.service.zenohservice import ZenohConfig, ZenohService, ZenohSessionPool +from dimos.protocol.service.zenohservice import ( + ZENOH_LOCAL_ROUTER_ENDPOINT, + ZENOH_ROUTER_ENDPOINT_ENV, + ZenohConfig, + ZenohService, + ZenohSessionPool, +) @pytest.fixture() @@ -33,6 +39,17 @@ def test_different_modes_produce_different_keys() -> None: assert peer.session_key != client.session_key +def test_default_config_uses_local_router_when_configured( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv(ZENOH_ROUTER_ENDPOINT_ENV, ZENOH_LOCAL_ROUTER_ENDPOINT) + + config = ZenohConfig() + + assert config.mode == "client" + assert config.connect == [ZENOH_LOCAL_ROUTER_ENDPOINT] + + def test_start_creates_session(session_pool) -> None: svc = ZenohService(session_pool=session_pool) svc.start() diff --git a/dimos/protocol/service/zenohservice.py b/dimos/protocol/service/zenohservice.py index 22bc364377..227df0c29b 100644 --- a/dimos/protocol/service/zenohservice.py +++ b/dimos/protocol/service/zenohservice.py @@ -15,6 +15,7 @@ from __future__ import annotations import json +import os import platform import socket import threading @@ -31,6 +32,10 @@ logger = setup_logger() +ZENOH_ROUTER_ENDPOINT_ENV = "DIMOS_ZENOH_ROUTER_ENDPOINT" +ZENOH_LOCAL_ROUTER_ENDPOINT = "tcp/127.0.0.1:7447" +ZENOH_LOCAL_ROUTER_LISTEN = "tcp/[::]:7447" + # Robot-side bridges (e.g. go2web) listen here so a remote dimos can dial in # when multicast discovery fails. Zenoh's own default port. ROBOT_ZENOH_PORT = 7447 @@ -43,14 +48,12 @@ LOOPBACK_INTERFACE = "lo0" if platform.system() == "Darwin" else "lo" -def _default_connect_endpoints() -> list[str]: - """Dial known robots directly instead of trusting multicast scouting. +def _robot_connect_endpoints() -> list[str]: + """Return explicit robot endpoints when multicast scouting is insufficient. Many APs filter multicast between WiFi clients, so a robot that is - perfectly reachable over TCP never answers a scout. When the session is - zenoh-transported and a robot IP is configured, it becomes an explicit - endpoint; scouting stays on for everything else. An IP carrying its own - ``:port`` is used as given. + reachable over TCP may never answer a scout. An IP carrying its own port is + used as given. """ from dimos.core.global_config import global_config @@ -67,6 +70,17 @@ def _default_connect_endpoints() -> list[str]: return out +def _default_mode() -> str: + return "client" if os.getenv(ZENOH_ROUTER_ENDPOINT_ENV) else "peer" + + +def _default_connect_endpoints() -> list[str]: + router_endpoint = os.getenv(ZENOH_ROUTER_ENDPOINT_ENV) + if router_endpoint: + return [router_endpoint] + return _robot_connect_endpoints() + + def _default_scouting() -> bool: from dimos.core.global_config import global_config @@ -98,10 +112,37 @@ def endpoint_addresses(endpoint: str) -> set[str]: return out +def _await_connect_endpoints( + session: zenoh.Session, + endpoints: list[str], + timeout: float, +) -> None: + pending = {endpoint: endpoint_addresses(endpoint) for endpoint in endpoints} + if not pending or timeout <= 0: + return + + deadline = time.monotonic() + timeout + while pending: + linked = {str(link.dst).rpartition("/")[2] for link in session.info.links()} + for endpoint in [item for item, addresses in pending.items() if addresses & linked]: + logger.debug("Zenoh linked", endpoint=endpoint) + del pending[endpoint] + if not pending: + return + if time.monotonic() >= deadline: + logger.warning( + "Zenoh endpoints not linked; continuing", + timeout=timeout, + endpoints=sorted(pending), + ) + return + time.sleep(_CONNECT_POLL_INTERVAL) + + class ZenohConfig(BaseConfig): - mode: str = "peer" + mode: str = Field(default_factory=_default_mode) connect: list[str] = Field(default_factory=_default_connect_endpoints) - listen: list[str] = [] + listen: list[str] = Field(default_factory=list) # Discover peers across the network. Off keeps discovery on loopback. scouting: bool = Field(default_factory=_default_scouting) # Seconds to block in start() waiting for `connect` endpoints to link. @@ -156,6 +197,44 @@ def close_all(self) -> None: default_session_pool = ZenohSessionPool() +class ZenohRouter: + def __init__( + self, + listen: str = ZENOH_LOCAL_ROUTER_LISTEN, + connect: list[str] | None = None, + ) -> None: + self._listen = listen + self._connect = _robot_connect_endpoints() if connect is None else list(connect) + self._session: zenoh.Session | None = None + + def start(self) -> None: + config = zenoh.Config() + config.insert_json5("mode", '"router"') + config.insert_json5("listen/endpoints", json.dumps([self._listen])) + if self._connect: + config.insert_json5("connect/endpoints", json.dumps(self._connect)) + if not _default_scouting(): + config.insert_json5("scouting/multicast/interface", json.dumps(LOOPBACK_INTERFACE)) + config.insert_json5("scouting/gossip/enabled", "false") + try: + self._session = zenoh.open(config) + _await_connect_endpoints( + self._session, + self._connect, + _default_connect_timeout(), + ) + logger.info("Local Zenoh router started", endpoint=self._listen) + except zenoh.ZError as exc: + if "Address already in use" not in str(exc): + raise + logger.info("Using existing local Zenoh router", endpoint=self._listen) + + def stop(self) -> None: + if self._session is not None: + self._session.close() + self._session = None + + class ZenohService(Service): config: ZenohConfig @@ -167,6 +246,9 @@ def __init__(self, *, session_pool: ZenohSessionPool | None = None, **kwargs: An self._session: zenoh.Session | None = None def start(self) -> None: + endpoint = os.getenv(ZENOH_ROUTER_ENDPOINT_ENV) + if endpoint and not self.config.model_fields_set: + self.config = ZenohConfig(mode="client", connect=[endpoint]) self._session = self._session_pool.acquire(self.config) self._await_connect(self._session) super().start() @@ -182,24 +264,11 @@ def _await_connect(self, session: zenoh.Session) -> None: Unreachable endpoints are a warning, not an error: one robot being down should not stop the rest of the graph from coming up. """ - pending = {ep: endpoint_addresses(ep) for ep in self.config.connect} - if not pending or self.config.connect_timeout <= 0: - return - deadline = time.monotonic() + self.config.connect_timeout - while pending: - linked = {str(link.dst).rpartition("/")[2] for link in session.info.links()} - for endpoint in [e for e, addrs in pending.items() if addrs & linked]: - logger.debug(f"Zenoh linked {endpoint}") - del pending[endpoint] - if not pending: - return - if time.monotonic() >= deadline: - logger.warning( - f"Zenoh endpoints not linked after {self.config.connect_timeout}s: " - f"{sorted(pending)} - continuing, published messages may be dropped" - ) - return - time.sleep(_CONNECT_POLL_INTERVAL) + _await_connect_endpoints( + session, + self.config.connect, + self.config.connect_timeout, + ) @property def session(self) -> zenoh.Session: diff --git a/dimos/robot/unitree/g1/blueprints/basic/unitree_g1_groot_wbc.py b/dimos/robot/unitree/g1/blueprints/basic/unitree_g1_groot_wbc.py index f9604c57d6..b19967d533 100644 --- a/dimos/robot/unitree/g1/blueprints/basic/unitree_g1_groot_wbc.py +++ b/dimos/robot/unitree/g1/blueprints/basic/unitree_g1_groot_wbc.py @@ -80,6 +80,7 @@ g1_urdf_joint_state, g1_urdf_static_robot, ) +from dimos.simulation.providers import SimulationRequest, load_simulation_provider from dimos.simulation.scene_assets.spec import ScenePackage from dimos.utils.data import LfsPath from dimos.visualization.rerun.scene_package import scene_package_static_entities @@ -94,6 +95,8 @@ _ROBOT_MESHDIR = LfsPath("g1_urdf/meshes") _adapter_address: str | Path +_using_simulation_provider = False +_provider_rerun_config: dict[str, Any] = {} _cmd_vel_topic = "/cmd_vel" if global_config.simulation else "/g1/cmd_vel" _MUJOCO_LIDAR_CAMERAS = ( "lidar_front_camera", @@ -265,19 +268,36 @@ def _precomposed_g1_scene(package: ScenePackage) -> Path | None: ) return candidate - # Sim backend: MuJoCo engine via SHM. - _backend, _adapter_address = _scene_mujoco_backend() + if global_config.simulation_provider: + _using_simulation_provider = True + _provider = load_simulation_provider(global_config.simulation_provider) + _binding = _provider.build( + SimulationRequest( + robot_model="unitree_g1", + model_path=_ROBOT_ONLY_MJCF_PATH, + mesh_dir=_ROBOT_MESHDIR, + scene_package=global_config.scene_package, + ) + ) + _backend = _binding.backend + _adapter_address = _binding.adapter_address + _adapter_type = _binding.adapter_type + _provider_rerun_config = _binding.rerun_config + else: + # Legacy in-repo MuJoCo backend. + _backend, _adapter_address = _scene_mujoco_backend() + _adapter_type = "sim_mujoco_g1" + # MujocoSimModule's ``odom`` Out is the sole producer of ``/odom`` # now - the coordinator no longer polls the whole-body adapter for # base pose (read_odom was dropped from the Protocol). autoconnect # maps ``(odom, PoseStamped)`` to ``/odom`` by default; no override. - _adapter_type = "sim_mujoco_g1" _tick_rate = 50.0 _auto_arm = True _auto_dry_run = False _default_ramp_seconds = 0.0 _decimation: int | None = 1 - _n_workers = 2 # sim: keep the default worker count + _n_workers = 12 if _using_simulation_provider else 2 _arm_holder = TaskConfig( name="servo_arms", type="servo", @@ -303,10 +323,9 @@ def _precomposed_g1_scene(package: ScenePackage) -> Path | None: ), MovementManager.blueprint(), ) - _remappings = [ - (VoxelGridMapper, "lidar", "pointcloud"), - (_G1GrootCoordinator, "twist_command", "cmd_vel"), - ] + _remappings = [(_G1GrootCoordinator, "twist_command", "cmd_vel")] + if not _using_simulation_provider: + _remappings.insert(0, (VoxelGridMapper, "lidar", "pointcloud")) else: from dimos.hardware.sensors.lidar.pointlio.module import PointLio from dimos.mapping.ray_tracing.module import RayTracingVoxelMap @@ -443,7 +462,8 @@ def _g1_real_costmap(grid: Any) -> Any: _static_rerun_entities: dict[str, Any] = { _G1_ROOT: g1_urdf_static_robot(root_path=_G1_ROOT), } -_static_rerun_entities.update(scene_package_static_entities(global_config.scene_package)) +if not _using_simulation_provider: + _static_rerun_entities.update(scene_package_static_entities(global_config.scene_package)) _rerun_config: dict[str, Any] = { "blueprint": _g1_groot_rerun_blueprint, @@ -476,6 +496,15 @@ def _g1_real_costmap(grid: Any) -> Any: "static": _static_rerun_entities, } +for _section in ("static", "visual_override", "max_hz"): + _rerun_config[_section] = { + **_rerun_config.get(_section, {}), + **_provider_rerun_config.get(_section, {}), + } +for _key, _value in _provider_rerun_config.items(): + if _key not in {"static", "visual_override", "max_hz"}: + _rerun_config[_key] = _value + if global_config.simulation != "mujoco": _rerun_config["visual_override"]["world/odometry"] = _g1_real_odometry_root _rerun_config["visual_override"]["world/global_costmap"] = _g1_real_costmap diff --git a/dimos/robot/unitree/go2/blueprints/basic/go2_platform.py b/dimos/robot/unitree/go2/blueprints/basic/go2_platform.py new file mode 100644 index 0000000000..d1ef72c734 --- /dev/null +++ b/dimos/robot/unitree/go2/blueprints/basic/go2_platform.py @@ -0,0 +1,55 @@ +# Copyright 2025-2026 Dimensional Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from functools import cache + +from dimos.core.coordination.blueprints import Blueprint +from dimos.core.global_config import global_config +from dimos.robot.unitree.go2.connection import GO2Connection +from dimos.simulation.providers import ( + SimulationBinding, + SimulationRequest, + load_simulation_provider, +) + + +def resolve_go2_platform() -> Blueprint: + if global_config.simulation in ("", "dimsim"): + return GO2Connection.blueprint() + return _resolve_simulation_binding().backend + + +def resolve_go2_rerun_config() -> dict[str, object]: + if global_config.simulation in ("", "dimsim"): + return {} + return _resolve_simulation_binding().rerun_config + + +@cache +def _resolve_simulation_binding() -> SimulationBinding: + if global_config.simulation != "mujoco": + raise ValueError("unitree-go2 only supports --simulation mujoco") + if not global_config.simulation_provider: + raise ValueError("unitree-go2 simulation requires --simulation-provider pimsim") + + provider = load_simulation_provider(global_config.simulation_provider) + return provider.build( + SimulationRequest( + robot_model="unitree_go2", + scene_package=global_config.scene_package, + ) + ) + + +__all__ = ["resolve_go2_platform", "resolve_go2_rerun_config"] diff --git a/dimos/robot/unitree/go2/blueprints/basic/test_go2_platform.py b/dimos/robot/unitree/go2/blueprints/basic/test_go2_platform.py new file mode 100644 index 0000000000..491696bde3 --- /dev/null +++ b/dimos/robot/unitree/go2/blueprints/basic/test_go2_platform.py @@ -0,0 +1,64 @@ +# Copyright 2025-2026 Dimensional Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from collections.abc import Iterator +from pathlib import Path + +import pytest + +from dimos.core.coordination.blueprints import Blueprint +from dimos.core.global_config import global_config +from dimos.robot.unitree.go2.blueprints.basic import go2_platform +from dimos.simulation.providers import SimulationBinding, SimulationRequest + + +@pytest.fixture +def pimsim_go2_config(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]: + monkeypatch.setattr(global_config, "simulation", "mujoco") + monkeypatch.setattr(global_config, "simulation_provider", "pimsim") + monkeypatch.setattr(global_config, "scene_package", "dimsim-apartment") + go2_platform._resolve_simulation_binding.cache_clear() + yield + go2_platform._resolve_simulation_binding.cache_clear() + + +def test_go2_platform_uses_requested_simulation_provider( + pimsim_go2_config: None, + mocker, +) -> None: + del pimsim_go2_config + backend = Blueprint(blueprints=()) + binding = SimulationBinding( + backend=backend, + adapter_type="pimsim_go2", + adapter_address=Path("/tmp/pimsim-go2"), + rerun_config={"static": {"world/scene": "scene"}}, + ) + provider = mocker.Mock() + provider.build.return_value = binding + load_provider = mocker.patch.object( + go2_platform, + "load_simulation_provider", + return_value=provider, + ) + + assert go2_platform.resolve_go2_platform() is backend + assert go2_platform.resolve_go2_rerun_config() == binding.rerun_config + load_provider.assert_called_once_with("pimsim") + provider.build.assert_called_once_with( + SimulationRequest( + robot_model="unitree_go2", + scene_package="dimsim-apartment", + ) + ) diff --git a/dimos/robot/unitree/go2/blueprints/basic/unitree_go2_basic.py b/dimos/robot/unitree/go2/blueprints/basic/unitree_go2_basic.py index c342c74f92..b850f654d7 100644 --- a/dimos/robot/unitree/go2/blueprints/basic/unitree_go2_basic.py +++ b/dimos/robot/unitree/go2/blueprints/basic/unitree_go2_basic.py @@ -18,7 +18,10 @@ from dimos.core.coordination.blueprints import autoconnect from dimos.core.global_config import global_config -from dimos.robot.unitree.go2.connection import GO2Connection +from dimos.robot.unitree.go2.blueprints.basic.go2_platform import ( + resolve_go2_platform, + resolve_go2_rerun_config, +) from dimos.visualization.vis_module import vis_module @@ -103,6 +106,16 @@ def _go2_rerun_blueprint() -> Any: }, } +_provider_rerun_config = resolve_go2_rerun_config() +for _section in ("static", "visual_override", "max_hz"): + rerun_config[_section] = { + **rerun_config.get(_section, {}), + **_provider_rerun_config.get(_section, {}), + } +for _key, _value in _provider_rerun_config.items(): + if _key not in {"static", "visual_override", "max_hz"}: + rerun_config[_key] = _value + _with_vis = autoconnect( vis_module( viewer_backend=global_config.viewer, @@ -114,7 +127,7 @@ def _go2_rerun_blueprint() -> Any: unitree_go2_basic = ( autoconnect( _with_vis, - GO2Connection.blueprint(), + resolve_go2_platform(), ).global_config(n_workers=4, robot_model="unitree_go2") # we temporarily disabled sensor timestamps # and are derriving all timestmaps upon reception diff --git a/dimos/simulation/adapters/whole_body/g1.py b/dimos/simulation/adapters/whole_body/g1.py index d29e1b6585..6f17a5a183 100644 --- a/dimos/simulation/adapters/whole_body/g1.py +++ b/dimos/simulation/adapters/whole_body/g1.py @@ -29,16 +29,16 @@ import time from typing import Any +from dimos.hardware.simulation.shared_memory import ( + ManipShmReader, + shm_key_from_path, +) from dimos.hardware.whole_body.spec import ( POS_STOP, IMUState, MotorCommand, MotorState, ) -from dimos.simulation.engines.mujoco_shm import ( - ManipShmReader, - shm_key_from_path, -) from dimos.utils.logging_config import setup_logger logger = setup_logger() diff --git a/dimos/simulation/engines/mujoco_shm.py b/dimos/simulation/engines/mujoco_shm.py index 7cb24e39a7..33489d8f81 100644 --- a/dimos/simulation/engines/mujoco_shm.py +++ b/dimos/simulation/engines/mujoco_shm.py @@ -12,479 +12,56 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Shared-memory buffers for sim-manipulator IPC. - -Layout for exchanging joint state and commands between ``MujocoSimModule`` -(which owns the physics engine) and ``ShmMujocoAdapter`` (which plugs into -ControlCoordinator). Modeled after ``dimos.simulation.mujoco.shared_memory`` -(the Go2 SHM pattern). - -Names are deterministic: both sides derive them from the resolved MJCF path, -so no name exchange over RPC is needed. The sim module creates the buffers -and signals ``ready``; the adapter attaches to them by name. -""" - -from __future__ import annotations - -from dataclasses import dataclass -import hashlib -from multiprocessing import resource_tracker -from multiprocessing.shared_memory import SharedMemory -from pathlib import Path -from typing import Any - -import numpy as np -from numpy.typing import NDArray - -from dimos.utils.logging_config import setup_logger - -logger = setup_logger() - -# Upper bound on joint count per sim. Manipulators use <=10; humanoids -# (Unitree G1: 29) push higher. 32 leaves headroom while keeping all -# per-joint buffers tiny (32 floats = 256 B). -MAX_JOINTS = 32 -_FLOAT_BYTES = 8 # float64 -_INT32_BYTES = 4 - -# IMU layout: quat (4) + gyro (3) + accel (3) = 10 floats. -_IMU_FLOATS = 10 - -_joint_array_size = MAX_JOINTS * _FLOAT_BYTES # float64 array - -# Element counts for control and sequence arrays. -_NUM_CTRL_FIELDS = 4 # [ready, stop, command_mode, num_joints] -_NUM_SEQ_COUNTERS = 12 # one per buffer type (manipulator + WB additions) - -# Buffer sizes (in bytes). -# Keys are short to stay under macOS PSHMNAMLEN (31 bytes). -_shm_sizes = { - # Manipulator-shared layout - "pos": _joint_array_size, - "vel": _joint_array_size, - "eff": _joint_array_size, - "pos_t": _joint_array_size, - "vel_t": _joint_array_size, - "grp": 2 * _FLOAT_BYTES, # [gripper_position, gripper_target] - # Whole-body additions (unused by manipulator path). - "imu": _IMU_FLOATS * _FLOAT_BYTES, # [w,x,y,z, gx,gy,gz, ax,ay,az] - "kp_t": _joint_array_size, # per-joint position-gain target - "kd_t": _joint_array_size, # per-joint velocity-gain target - "tau_t": _joint_array_size, # per-joint feedforward torque - # Bookkeeping - "seq": _NUM_SEQ_COUNTERS * _FLOAT_BYTES, # int64 counters - "ctl": _NUM_CTRL_FIELDS * _INT32_BYTES, # [ready, stop, command_mode, num_joints] -} - -# Sequence counter indices. -SEQ_POSITIONS = 0 -SEQ_VELOCITIES = 1 -SEQ_EFFORTS = 2 -SEQ_POSITION_CMD = 3 -SEQ_VELOCITY_CMD = 4 -SEQ_GRIPPER_STATE = 5 -SEQ_GRIPPER_CMD = 6 -# Whole-body additions -SEQ_IMU = 7 -SEQ_KP_CMD = 8 -SEQ_KD_CMD = 9 -SEQ_TAU_CMD = 10 - -# Control indices. -CTRL_READY = 0 -CTRL_STOP = 1 -CTRL_COMMAND_MODE = 2 -CTRL_NUM_JOINTS = 3 - -# Command modes. -CMD_MODE_POSITION = 0 -CMD_MODE_VELOCITY = 1 -# Whole-body PD-with-feedforward: ctrl = kp*(q_t - q) + kd*(0 - dq) + tau_t. -# Per-step kp/kd lets a policy retune gains online if it wants to. -CMD_MODE_PD_TAU = 2 - -_NAME_PREFIX = "dmjm" - - -def shm_key_from_path(config_path: Path | str) -> str: - """Derive a deterministic short key from an MJCF path. - - Both sim module and adapter compute the same key from the same path, - so SHM buffer names can be agreed upon without an RPC round-trip. - """ - resolved = str(Path(config_path).expanduser().resolve()) - return hashlib.md5(resolved.encode("utf-8")).hexdigest()[:12] - - -def _buffer_name(key: str, buffer: str) -> str: - return f"{_NAME_PREFIX}_{key}_{buffer}" - - -def _unregister(shm: SharedMemory) -> SharedMemory: - """Detach ``shm`` from ``resource_tracker`` to silence spurious warnings. - - Same technique as ``dimos.simulation.mujoco.shared_memory._unregister``. - """ - try: - resource_tracker.unregister(shm._name, "shared_memory") # type: ignore[attr-defined] - except Exception: - pass - return shm - - -@dataclass(frozen=True) -class ManipShmSet: - """Frozen set of named SharedMemory buffers for sim <-> adapter IPC. - - Despite the name (kept for backward compat with existing manipulator - consumers), the layout now also covers whole-body needs: IMU, per-joint - PD gain commands, and per-joint feedforward torque commands. The - extra buffers are unused by the manipulator path. - """ - - pos: SharedMemory - vel: SharedMemory - eff: SharedMemory - pos_t: SharedMemory - vel_t: SharedMemory - grp: SharedMemory - # Whole-body additions - imu: SharedMemory - kp_t: SharedMemory - kd_t: SharedMemory - tau_t: SharedMemory - # Bookkeeping - seq: SharedMemory - ctl: SharedMemory - - @classmethod - def create(cls, key: str) -> ManipShmSet: - """Create new SHM buffers with deterministic names derived from *key*""" - buffers: dict[str, SharedMemory] = {} - for buffer_name, size in _shm_sizes.items(): - name = _buffer_name(key, buffer_name) - try: - stale = _unregister(SharedMemory(name=name)) - stale.close() - try: - stale.unlink() - logger.info("ManipShmSet: unlinked stale SHM", name=name) - except FileNotFoundError: - pass - except FileNotFoundError: - pass - buffers[buffer_name] = SharedMemory(create=True, size=size, name=name) - return cls(**buffers) - - @classmethod - def attach(cls, key: str) -> ManipShmSet: - """Attach to existing SHM buffers created by the sim side.""" - buffers: dict[str, SharedMemory] = {} - for buffer_name in _shm_sizes: - name = _buffer_name(key, buffer_name) - buffers[buffer_name] = _unregister(SharedMemory(name=name)) - return cls(**buffers) - - def as_list(self) -> list[SharedMemory]: - return [getattr(self, k) for k in _shm_sizes] - - -class ManipShmWriter: - """Sim-side handle: writes joint state, reads command targets. - Owned by ``MujocoSimModule``. Creates the SHM buffers on init and - unlinks them on cleanup. - """ - - shm: ManipShmSet - - def __init__(self, key: str) -> None: - self.shm = ManipShmSet.create(key) - self._last_pos_cmd_seq = 0 - self._last_vel_cmd_seq = 0 - self._last_gripper_cmd_seq = 0 - self._last_kp_cmd_seq = 0 - self._last_kd_cmd_seq = 0 - self._last_tau_cmd_seq = 0 - # Zero everything. - for buf in self.shm.as_list(): - np.ndarray((buf.size,), dtype=np.uint8, buffer=buf.buf)[:] = 0 - - def write_joint_state( - self, - positions: list[float], - velocities: list[float], - efforts: list[float], - ) -> None: - n = min(len(positions), MAX_JOINTS) - pos_arr = self._array(self.shm.pos, MAX_JOINTS, np.float64) - vel_arr = self._array(self.shm.vel, MAX_JOINTS, np.float64) - eff_arr = self._array(self.shm.eff, MAX_JOINTS, np.float64) - pos_arr[:n] = positions[:n] - vel_arr[:n] = velocities[:n] - eff_arr[:n] = efforts[:n] - self._increment_seq(SEQ_POSITIONS) - self._increment_seq(SEQ_VELOCITIES) - self._increment_seq(SEQ_EFFORTS) - - def write_gripper_state(self, position: float) -> None: - arr = self._array(self.shm.grp, 2, np.float64) - arr[0] = position - self._increment_seq(SEQ_GRIPPER_STATE) - - def read_position_command(self, num_joints: int) -> NDArray[np.float64] | None: - """Return a copy of position targets if a new command arrived since last call.""" - seq = self._get_seq(SEQ_POSITION_CMD) - if seq <= self._last_pos_cmd_seq: - return None - self._last_pos_cmd_seq = seq - arr = self._array(self.shm.pos_t, MAX_JOINTS, np.float64) - result: NDArray[np.float64] = arr[:num_joints].copy() - return result - - def read_velocity_command(self, num_joints: int) -> NDArray[np.float64] | None: - seq = self._get_seq(SEQ_VELOCITY_CMD) - if seq <= self._last_vel_cmd_seq: - return None - self._last_vel_cmd_seq = seq - arr = self._array(self.shm.vel_t, MAX_JOINTS, np.float64) - result: NDArray[np.float64] = arr[:num_joints].copy() - return result - - def read_gripper_command(self) -> float | None: - seq = self._get_seq(SEQ_GRIPPER_CMD) - if seq <= self._last_gripper_cmd_seq: - return None - self._last_gripper_cmd_seq = seq - arr = self._array(self.shm.grp, 2, np.float64) - return float(arr[1]) - - def read_command_mode(self) -> int: - return int(self._control()[CTRL_COMMAND_MODE]) - - # Whole-body additions - - def write_imu( - self, - quaternion: tuple[float, float, float, float], - gyroscope: tuple[float, float, float], - accelerometer: tuple[float, float, float], - ) -> None: - """Write IMU sample. Quaternion is (w, x, y, z).""" - arr = self._array(self.shm.imu, _IMU_FLOATS, np.float64) - arr[0:4] = quaternion - arr[4:7] = gyroscope - arr[7:10] = accelerometer - self._increment_seq(SEQ_IMU) - - def read_kp_command(self, num_joints: int) -> NDArray[np.float64] | None: - """Per-joint position-gain target if a new command landed since last call.""" - seq = self._get_seq(SEQ_KP_CMD) - if seq <= self._last_kp_cmd_seq: - return None - self._last_kp_cmd_seq = seq - arr = self._array(self.shm.kp_t, MAX_JOINTS, np.float64) - return arr[:num_joints].copy() - - def read_kd_command(self, num_joints: int) -> NDArray[np.float64] | None: - seq = self._get_seq(SEQ_KD_CMD) - if seq <= self._last_kd_cmd_seq: - return None - self._last_kd_cmd_seq = seq - arr = self._array(self.shm.kd_t, MAX_JOINTS, np.float64) - return arr[:num_joints].copy() - - def read_tau_command(self, num_joints: int) -> NDArray[np.float64] | None: - """Per-joint feedforward torque if a new command landed since last call.""" - seq = self._get_seq(SEQ_TAU_CMD) - if seq <= self._last_tau_cmd_seq: - return None - self._last_tau_cmd_seq = seq - arr = self._array(self.shm.tau_t, MAX_JOINTS, np.float64) - return arr[:num_joints].copy() - - def signal_ready(self, num_joints: int) -> None: - ctrl = self._control() - ctrl[CTRL_NUM_JOINTS] = num_joints - ctrl[CTRL_READY] = 1 - - def signal_stop(self) -> None: - self._control()[CTRL_STOP] = 1 - - def should_stop(self) -> bool: - return bool(self._control()[CTRL_STOP] == 1) - - def cleanup(self) -> None: - for shm in self.shm.as_list(): - try: - shm.close() - except FileNotFoundError: - pass # already detached - except OSError as exc: - logger.warning("SHM close failed", name=shm.name, error=str(exc)) - try: - shm.unlink() - except FileNotFoundError: - pass # already unlinked (e.g. cleanup called twice) - except OSError as exc: - logger.warning("SHM unlink failed", name=shm.name, error=str(exc)) - - def _array(self, buf: SharedMemory, n: int, dtype: Any) -> NDArray[Any]: - return np.ndarray((n,), dtype=dtype, buffer=buf.buf) - - def _control(self) -> NDArray[np.int32]: - return np.ndarray((_NUM_CTRL_FIELDS,), dtype=np.int32, buffer=self.shm.ctl.buf) - - def _increment_seq(self, index: int) -> None: - seq_arr = np.ndarray((_NUM_SEQ_COUNTERS,), dtype=np.int64, buffer=self.shm.seq.buf) - seq_arr[index] += 1 - - def _get_seq(self, index: int) -> int: - seq_arr = np.ndarray((_NUM_SEQ_COUNTERS,), dtype=np.int64, buffer=self.shm.seq.buf) - return int(seq_arr[index]) - - -class ManipShmReader: - """Adapter-side handle: reads joint state, writes command targets. - - Owned by ``ShmMujocoAdapter``. Attaches to existing buffers created by - the sim module; does not unlink them on cleanup. - """ - - shm: ManipShmSet - - def __init__(self, key: str) -> None: - self.shm = ManipShmSet.attach(key) - - def read_positions(self, num_joints: int) -> list[float]: - arr = np.ndarray((MAX_JOINTS,), dtype=np.float64, buffer=self.shm.pos.buf) - return [float(x) for x in arr[:num_joints]] - - def read_velocities(self, num_joints: int) -> list[float]: - arr = np.ndarray((MAX_JOINTS,), dtype=np.float64, buffer=self.shm.vel.buf) - return [float(x) for x in arr[:num_joints]] - - def read_efforts(self, num_joints: int) -> list[float]: - arr = np.ndarray((MAX_JOINTS,), dtype=np.float64, buffer=self.shm.eff.buf) - return [float(x) for x in arr[:num_joints]] - - def read_gripper_position(self) -> float: - arr = np.ndarray((2,), dtype=np.float64, buffer=self.shm.grp.buf) - return float(arr[0]) - - def write_position_command(self, positions: list[float]) -> None: - n = min(len(positions), MAX_JOINTS) - arr = np.ndarray((MAX_JOINTS,), dtype=np.float64, buffer=self.shm.pos_t.buf) - arr[:n] = positions[:n] - self._set_command_mode(CMD_MODE_POSITION) - self._increment_seq(SEQ_POSITION_CMD) - - def write_velocity_command(self, velocities: list[float]) -> None: - n = min(len(velocities), MAX_JOINTS) - arr = np.ndarray((MAX_JOINTS,), dtype=np.float64, buffer=self.shm.vel_t.buf) - arr[:n] = velocities[:n] - self._set_command_mode(CMD_MODE_VELOCITY) - self._increment_seq(SEQ_VELOCITY_CMD) - - def write_gripper_command(self, position: float) -> None: - arr = np.ndarray((2,), dtype=np.float64, buffer=self.shm.grp.buf) - arr[1] = position - self._increment_seq(SEQ_GRIPPER_CMD) - - # Whole-body additions - - def read_imu( - self, - ) -> tuple[ - tuple[float, float, float, float], - tuple[float, float, float], - tuple[float, float, float], - ]: - """Read IMU sample: ((qw, qx, qy, qz), (gx, gy, gz), (ax, ay, az)).""" - arr = np.ndarray((_IMU_FLOATS,), dtype=np.float64, buffer=self.shm.imu.buf) - return ( - (float(arr[0]), float(arr[1]), float(arr[2]), float(arr[3])), - (float(arr[4]), float(arr[5]), float(arr[6])), - (float(arr[7]), float(arr[8]), float(arr[9])), - ) - - def write_kp_command(self, kp: list[float]) -> None: - """Per-joint position-gain target. Switches command mode to PD+tau.""" - n = min(len(kp), MAX_JOINTS) - arr = np.ndarray((MAX_JOINTS,), dtype=np.float64, buffer=self.shm.kp_t.buf) - arr[:n] = kp[:n] - self._set_command_mode(CMD_MODE_PD_TAU) - self._increment_seq(SEQ_KP_CMD) - - def write_kd_command(self, kd: list[float]) -> None: - n = min(len(kd), MAX_JOINTS) - arr = np.ndarray((MAX_JOINTS,), dtype=np.float64, buffer=self.shm.kd_t.buf) - arr[:n] = kd[:n] - self._set_command_mode(CMD_MODE_PD_TAU) - self._increment_seq(SEQ_KD_CMD) - - def write_tau_command(self, tau: list[float]) -> None: - """Per-joint feedforward torque, applied on top of PD.""" - n = min(len(tau), MAX_JOINTS) - arr = np.ndarray((MAX_JOINTS,), dtype=np.float64, buffer=self.shm.tau_t.buf) - arr[:n] = tau[:n] - self._set_command_mode(CMD_MODE_PD_TAU) - self._increment_seq(SEQ_TAU_CMD) - - def write_pd_tau_command( - self, - positions: list[float], - kp: list[float], - kd: list[float], - tau: list[float], - ) -> None: - """Write a whole-body PD+tau command without transient mode flips. - - The sim engine runs in a different process, so setting position mode - first and PD mode later creates a small but real race. Write all arrays, - publish PD mode once, then bump the sequence counters. - """ - n_pos = min(len(positions), MAX_JOINTS) - n_kp = min(len(kp), MAX_JOINTS) - n_kd = min(len(kd), MAX_JOINTS) - n_tau = min(len(tau), MAX_JOINTS) - np.ndarray((MAX_JOINTS,), dtype=np.float64, buffer=self.shm.pos_t.buf)[:n_pos] = positions[ - :n_pos - ] - np.ndarray((MAX_JOINTS,), dtype=np.float64, buffer=self.shm.kp_t.buf)[:n_kp] = kp[:n_kp] - np.ndarray((MAX_JOINTS,), dtype=np.float64, buffer=self.shm.kd_t.buf)[:n_kd] = kd[:n_kd] - np.ndarray((MAX_JOINTS,), dtype=np.float64, buffer=self.shm.tau_t.buf)[:n_tau] = tau[:n_tau] - self._set_command_mode(CMD_MODE_PD_TAU) - self._increment_seq(SEQ_KP_CMD) - self._increment_seq(SEQ_KD_CMD) - self._increment_seq(SEQ_TAU_CMD) - # Position is the engine-side trigger for latching a new PD target, - # so publish it last after gains/torque are visible. - self._increment_seq(SEQ_POSITION_CMD) - - def is_ready(self) -> bool: - return bool(self._control()[CTRL_READY] == 1) - - def num_joints(self) -> int: - return int(self._control()[CTRL_NUM_JOINTS]) - - def signal_stop(self) -> None: - self._control()[CTRL_STOP] = 1 - - def cleanup(self) -> None: - for shm in self.shm.as_list(): - try: - shm.close() - except FileNotFoundError: - pass # already detached - except OSError as exc: - logger.warning("SHM close failed", name=shm.name, error=str(exc)) - - def _control(self) -> NDArray[np.int32]: - return np.ndarray((_NUM_CTRL_FIELDS,), dtype=np.int32, buffer=self.shm.ctl.buf) - - def _set_command_mode(self, mode: int) -> None: - self._control()[CTRL_COMMAND_MODE] = mode - - def _increment_seq(self, index: int) -> None: - seq_arr = np.ndarray((_NUM_SEQ_COUNTERS,), dtype=np.int64, buffer=self.shm.seq.buf) - seq_arr[index] += 1 +"""Compatibility imports for the neutral simulation shared-memory boundary.""" + +from dimos.hardware.simulation.shared_memory import ( + CMD_MODE_PD_TAU, + CMD_MODE_POSITION, + CMD_MODE_VELOCITY, + CTRL_COMMAND_MODE, + CTRL_NUM_JOINTS, + CTRL_READY, + CTRL_STOP, + MAX_JOINTS, + SEQ_EFFORTS, + SEQ_GRIPPER_CMD, + SEQ_GRIPPER_STATE, + SEQ_IMU, + SEQ_KD_CMD, + SEQ_KP_CMD, + SEQ_POSITION_CMD, + SEQ_POSITIONS, + SEQ_TAU_CMD, + SEQ_VELOCITIES, + SEQ_VELOCITY_CMD, + ManipShmReader, + ManipShmSet, + ManipShmWriter, + shm_key_from_path, +) + +__all__ = [ + "CMD_MODE_PD_TAU", + "CMD_MODE_POSITION", + "CMD_MODE_VELOCITY", + "CTRL_COMMAND_MODE", + "CTRL_NUM_JOINTS", + "CTRL_READY", + "CTRL_STOP", + "MAX_JOINTS", + "SEQ_EFFORTS", + "SEQ_GRIPPER_CMD", + "SEQ_GRIPPER_STATE", + "SEQ_IMU", + "SEQ_KD_CMD", + "SEQ_KP_CMD", + "SEQ_POSITIONS", + "SEQ_POSITION_CMD", + "SEQ_TAU_CMD", + "SEQ_VELOCITIES", + "SEQ_VELOCITY_CMD", + "ManipShmReader", + "ManipShmSet", + "ManipShmWriter", + "shm_key_from_path", +] diff --git a/dimos/simulation/engines/mujoco_sim_module.py b/dimos/simulation/engines/mujoco_sim_module.py index 41151f722d..da468186ca 100644 --- a/dimos/simulation/engines/mujoco_sim_module.py +++ b/dimos/simulation/engines/mujoco_sim_module.py @@ -43,6 +43,11 @@ from dimos.core.module import Module, ModuleConfig from dimos.core.stream import Out from dimos.hardware.sensors.camera.spec import DepthCameraConfig, DepthCameraHardware +from dimos.hardware.simulation.shared_memory import ( + CMD_MODE_PD_TAU, + ManipShmWriter, + shm_key_from_path, +) from dimos.msgs.geometry_msgs.PoseStamped import PoseStamped from dimos.msgs.geometry_msgs.Quaternion import Quaternion from dimos.msgs.geometry_msgs.Transform import Transform @@ -59,11 +64,6 @@ MujocoEngine, RaycastLidarConfig, ) -from dimos.simulation.engines.mujoco_shm import ( - CMD_MODE_PD_TAU, - ManipShmWriter, - shm_key_from_path, -) from dimos.simulation.engines.robot_sim_binding import RobotSimSpec from dimos.simulation.mujoco.constants import LIDAR_RESOLUTION, MAX_HEIGHT, MAX_RANGE, MIN_RANGE from dimos.spec import perception diff --git a/dimos/simulation/providers.py b/dimos/simulation/providers.py new file mode 100644 index 0000000000..32b6a78bf2 --- /dev/null +++ b/dimos/simulation/providers.py @@ -0,0 +1,66 @@ +# Copyright 2025-2026 Dimensional Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from __future__ import annotations + +from dataclasses import dataclass, field +import importlib.metadata as importlib_metadata +from pathlib import Path +from typing import Any, Protocol, runtime_checkable + +from dimos.core.coordination.blueprints import Blueprint +from dimos.msgs.geometry_msgs.PoseStamped import PoseStamped + +ENTRY_POINT_GROUP = "dimos.simulation.providers" + + +@dataclass(frozen=True) +class SimulationRequest: + robot_model: str + model_path: str | Path | None = None + mesh_dir: str | Path | None = None + scene_package: str | Path | None = None + + +@dataclass(frozen=True) +class SimulationBinding: + backend: Blueprint + adapter_type: str + adapter_address: str | Path + rerun_config: dict[str, Any] = field(default_factory=dict) + robot_base_pose: PoseStamped = field(default_factory=PoseStamped) + + +@runtime_checkable +class SimulationProvider(Protocol): + def build(self, request: SimulationRequest) -> SimulationBinding: ... + + +def load_simulation_provider(name: str) -> SimulationProvider: + matches = list(importlib_metadata.entry_points(group=ENTRY_POINT_GROUP, name=name)) + if not matches: + available = sorted( + entry_point.name + for entry_point in importlib_metadata.entry_points(group=ENTRY_POINT_GROUP) + ) + suffix = f" Available providers: {', '.join(available)}." if available else "" + raise ValueError(f"Simulation provider {name!r} is not installed.{suffix}") + if len(matches) > 1: + raise ValueError(f"Simulation provider {name!r} is registered more than once") + provider = matches[0].load() + if not isinstance(provider, SimulationProvider): + raise TypeError( + f"Simulation provider {name!r} must implement SimulationProvider, got {provider!r}" + ) + return provider diff --git a/dimos/simulation/test_providers.py b/dimos/simulation/test_providers.py new file mode 100644 index 0000000000..8d2cd4c6ed --- /dev/null +++ b/dimos/simulation/test_providers.py @@ -0,0 +1,42 @@ +# Copyright 2025-2026 Dimensional Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from typing import Any + +import pytest + +from dimos.simulation import providers +from dimos.simulation.providers import SimulationBinding, SimulationRequest + + +class _Provider: + def build(self, request: SimulationRequest) -> SimulationBinding: + raise NotImplementedError + + +class _EntryPoint: + name = "test" + + def load(self) -> Any: + return _Provider() + + +def test_load_external_simulation_provider(monkeypatch: pytest.MonkeyPatch) -> None: + def entry_points(*, group: str, name: str | None = None) -> list[_EntryPoint]: + assert group == providers.ENTRY_POINT_GROUP + return [_EntryPoint()] if name in (None, "test") else [] + + monkeypatch.setattr(providers.importlib_metadata, "entry_points", entry_points) + + assert isinstance(providers.load_simulation_provider("test"), _Provider) diff --git a/dimos/visualization/rerun/bridge.py b/dimos/visualization/rerun/bridge.py index 25f16e0271..8fba9641a1 100644 --- a/dimos/visualization/rerun/bridge.py +++ b/dimos/visualization/rerun/bridge.py @@ -223,6 +223,7 @@ class Config(ModuleConfig): visual_override: dict[Glob | str, VisualOverride | None] = field(default_factory=dict) static: dict[str, Callable[[Any], Any]] = field(default_factory=dict) max_hz: dict[str, float] = field(default_factory=dict) + latest_state: set[str] = field(default_factory=set) entity_prefix: str = "world" topic_to_entity: Callable[[Any], str] | None = None @@ -352,15 +353,20 @@ def _on_message(self, msg: Any, topic: Any) -> None: # TFMessage for example returns list of (entity_path, archetype) tuples if is_rerun_multi(rerun_data): for path, archetype in rerun_data: - rr.log(path, archetype) + rr.log(path, archetype, static=path in self.config.latest_state) else: - rr.log(entity_path, cast("Archetype", rerun_data)) + latest_state = entity_path in self.config.latest_state + rr.log(entity_path, cast("Archetype", rerun_data), static=latest_state) # if source msg carries a frame_id, attach the entity to that TF frame # should skip if archetype is a Transform3D if not isinstance(rerun_data, rr.Transform3D): frame_id = getattr(msg, "frame_id", None) if frame_id and self._frame_attached.get(entity_path) != frame_id: - rr.log(entity_path, rr.Transform3D(parent_frame=f"tf#/{frame_id}")) + rr.log( + entity_path, + rr.Transform3D(parent_frame=f"tf#/{frame_id}"), + static=latest_state, + ) self._frame_attached[entity_path] = frame_id @rpc diff --git a/dimos/visualization/rerun/test_detection3d_bridge.py b/dimos/visualization/rerun/test_detection3d_bridge.py index 8385c6dfbf..46f8b4b86f 100644 --- a/dimos/visualization/rerun/test_detection3d_bridge.py +++ b/dimos/visualization/rerun/test_detection3d_bridge.py @@ -77,3 +77,18 @@ def test_detection3darray_bridge_attaches_topic_entity_to_message_frame() -> Non transform = mock_log.call_args_list[1].args[1] assert isinstance(transform, rr.Transform3D) assert transform.parent_frame.as_arrow_array().to_pylist() == ["tf#/world"] + + +def test_latest_state_entity_overwrites_data_and_frame_attachment() -> None: + entity_path = "world/marker_detection/detections" + bridge = RerunBridgeModule(latest_state={entity_path}) + bridge._min_intervals = {} + + try: + with patch("rerun.log") as mock_log: + bridge._on_message(_detection_array(), Topic("/marker_detection/detections")) + finally: + bridge.stop() + + assert mock_log.call_count == 2 + assert all(call.kwargs == {"static": True} for call in mock_log.call_args_list) diff --git a/docs/usage/transports/index.md b/docs/usage/transports/index.md index f27734fed9..27e34d93f3 100644 --- a/docs/usage/transports/index.md +++ b/docs/usage/transports/index.md @@ -351,6 +351,8 @@ Use Zenoh when: At the stream level, the transport wrappers are `ZenohTransport` and `pZenohTransport`. Install, defaults, and CLI versus environment overrides are in the [Zenoh quickstart](#zenoh-quickstart) above. +For a local `dimos run`, the coordinator starts or reuses one Zenoh router listening on port `7447`. Python workers, native modules, and local coordinator clients connect to it through `127.0.0.1`. This avoids an all-to-all peer mesh and keeps a local run independent of Wi-Fi, VPN, and Docker interface changes. The router remains reachable on other interfaces for explicitly configured remote participants. + Performance note: zenoh's session-to-session path (modules in different processes, the common case) benchmarks faster than LCM for small messages and for >=2MiB ones. Delivery *within* one shared session (co-located modules in one worker) is its slow path for 256KiB-1MiB messages (a few GiB/s); pin shared memory transports for heavy co-located streams. The benchmark has both cases (`Zenoh` = shared session, `ZenohPeers` = separate sessions). The Rerun bridge also follows the global transport. When `transport=zenoh`, the bridge listens on Zenoh and on LCM for TF data. diff --git a/native/rust/dimos-module/src/zenoh.rs b/native/rust/dimos-module/src/zenoh.rs index 9ff5209793..198ca89ce5 100644 --- a/native/rust/dimos-module/src/zenoh.rs +++ b/native/rust/dimos-module/src/zenoh.rs @@ -63,9 +63,16 @@ pub struct ZenohTransport { impl ZenohTransport { pub async fn new() -> io::Result { - let session = ::zenoh::open(::zenoh::Config::default()) - .await - .map_err(to_io)?; + let mut config = ::zenoh::Config::default(); + if let Ok(endpoint) = std::env::var("DIMOS_ZENOH_ROUTER_ENDPOINT") { + config.insert_json5("mode", r#""client""#).map_err(to_io)?; + let endpoints = serde_json::to_string(&[endpoint]) + .map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?; + config + .insert_json5("connect/endpoints", &endpoints) + .map_err(to_io)?; + } + let session = ::zenoh::open(config).await.map_err(to_io)?; Ok(Self { session, qos: OnceLock::new(),