From a1b0d2c93321e8fff22a8dff74287b14c2d8ba7b Mon Sep 17 00:00:00 2001 From: Matija Kecman Date: Wed, 1 Jul 2026 03:20:32 -0700 Subject: [PATCH] Add a Viewer-to-Sim Snapshot channel and convert relevant messages. The ViewerEndpoint can now send Snapshot messages to the SimEndpoint. The SimEndpoint has a new method, `get_viewer_snapshots`, to retrieve these. Thsi channel will be used for UI state needed for the simulation (e.g., `MjOption` and `StepControl`). PiperOrigin-RevId: 940981799 Change-Id: Ib9bd6ad848cd5479bc49492a34fddd82569cb09f --- .../mujoco/experimental/studio/endpoints.py | 20 +++++++++++++++++-- python/mujoco/experimental/studio/messages.py | 8 ++++---- 2 files changed, 22 insertions(+), 6 deletions(-) diff --git a/python/mujoco/experimental/studio/endpoints.py b/python/mujoco/experimental/studio/endpoints.py index edcbe2ce..35e6b207 100644 --- a/python/mujoco/experimental/studio/endpoints.py +++ b/python/mujoco/experimental/studio/endpoints.py @@ -29,10 +29,12 @@ class ViewerEndpoint: s2v_snapshot: vp.SnapshotChannel, s2v_events: vp.EventChannel, v2s_events: vp.EventChannel, + v2s_snapshot: vp.SnapshotChannel, ): self._s2v_snapshot = s2v_snapshot self._s2v_events = s2v_events self._v2s_events = v2s_events + self._v2s_snapshot = v2s_snapshot self._is_closed = False def send_to_sim(self, message: vp.Message) -> None: @@ -41,9 +43,10 @@ class ViewerEndpoint: if isinstance(message, vp.Event): self._v2s_events.put(message) elif isinstance(message, vp.Snapshot): - raise TypeError('viewer cannot produce snapshots') + self._v2s_snapshot.put(message) else: - raise TypeError(f'expected Event, got {type(message).__name__}') + name = type(message).__name__ + raise TypeError(f'expected Event or Snapshot, got {name}') def get_sim_events(self) -> list[vp.Event]: """Returns all pending events from the simulation.""" @@ -62,6 +65,7 @@ class ViewerEndpoint: if not self._is_closed: self._is_closed = True self._v2s_events.close() + self._v2s_snapshot.close() self._s2v_events.close() self._s2v_snapshot.close() @@ -79,10 +83,12 @@ class SimEndpoint: s2v_snapshot: vp.SnapshotChannel, s2v_events: vp.EventChannel, v2s_events: vp.EventChannel, + v2s_snapshot: vp.SnapshotChannel, ): self._s2v_snapshot = s2v_snapshot self._s2v_events = s2v_events self._v2s_events = v2s_events + self._v2s_snapshot = v2s_snapshot self._is_closed = False def send_to_viewer(self, message: vp.Message) -> None: @@ -103,6 +109,12 @@ class SimEndpoint: raise RuntimeError('SimEndpoint is closed') return self._v2s_events.get() + def get_viewer_snapshots(self) -> list[vp.Snapshot]: + """Returns all pending latest snapshots from the viewer, one per type.""" + if self._is_closed: + raise RuntimeError('SimEndpoint is closed') + return self._v2s_snapshot.get() + def close(self) -> None: """Close all channels owned by this endpoint.""" if not self._is_closed: @@ -110,6 +122,7 @@ class SimEndpoint: self._s2v_snapshot.close() self._s2v_events.close() self._v2s_events.close() + self._v2s_snapshot.close() def make_endpoints( @@ -117,16 +130,19 @@ def make_endpoints( s2v_snapshot: vp.SnapshotChannel, s2v_events: vp.EventChannel, v2s_events: vp.EventChannel, + v2s_snapshot: vp.SnapshotChannel, ) -> tuple[ViewerEndpoint, SimEndpoint]: """Returns viewer and simulation endpoints.""" viewer_endpoint = ViewerEndpoint( s2v_snapshot=s2v_snapshot, s2v_events=s2v_events, v2s_events=v2s_events, + v2s_snapshot=v2s_snapshot, ) sim_endpoint = SimEndpoint( s2v_snapshot=s2v_snapshot, s2v_events=s2v_events, v2s_events=v2s_events, + v2s_snapshot=v2s_snapshot, ) return viewer_endpoint, sim_endpoint diff --git a/python/mujoco/experimental/studio/messages.py b/python/mujoco/experimental/studio/messages.py index d36de578..62d29f2e 100644 --- a/python/mujoco/experimental/studio/messages.py +++ b/python/mujoco/experimental/studio/messages.py @@ -113,8 +113,8 @@ class ModelEvent(Event): @dataclasses.dataclass(frozen=True) -class OptionEvent(Event): - """An event that transports updated MuJoCo model options.""" +class MjOptionSnapshot(Snapshot): + """A snapshot sending mjOption state from viewer to sim each frame.""" opt: mujoco.MjOption @@ -128,8 +128,8 @@ class PerturbEvent(Event): @dataclasses.dataclass(frozen=True) -class StepControlEvent(Event): - """An event carrying the full step control state from the viewer to the sim.""" +class StepControlSnapshot(Snapshot): + """A snapshot sending step control state from viewer to sim each frame.""" pause_state: sim.PauseState speed: float