Refactor Studio customization and message dispatch to use decorator-based handlers

This change replaces fixed callback protocols (e.g., ViewerGuiHook, ViewerUpdateHook, SimEventHandler) with a general-purpose, priority-based message and event handling mechanism.

Key changes:
- Handler decorator and registry: Introduced the `@messages.handler(priority=...)` decorator and `HandlerRegistry` (`handler_registry.py`). Methods marked as handlers are automatically discovered and dispatched by priority (CRITICAL, USER, LIBRARY, INTERNAL) or method resolution order.
- Local lifecycle events: Added `ViewerAppInitEvent`, `BuildGuiEvent`, and `UpdateEvent` to `messages.py`. Custom GUI rendering and per-frame update logic can now be implemented as standard event handlers without needing separate interface protocols.
- Streamlined launch and app APIs: Replaced individual hook and handler arguments in `launch_passive`, `ViewerApp`, and `ViewerHandle` with unified `viewer_handlers` and `sim_handlers` lists.
- Module restructuring: Extracted simulation-side message handling and `ViewerHandle` from `sim_app.py` into a dedicated `viewer_handle.py` module, removing `sim_app.py`.
- Sample updates: Migrated existing examples (such as `implot.py`) to use the new handler pattern and lifecycle events.

PiperOrigin-RevId: 941729322
Change-Id: I93e7c0edf0a13a8dc854f9e2083451f7c25a1825
This commit is contained in:
Matija Kecman
2026-07-02 09:14:02 -07:00
committed by Copybara-Service
parent 1ca64b441b
commit 0dfa4b509a
7 changed files with 507 additions and 358 deletions
+73 -104
View File
@@ -14,10 +14,12 @@
"""Viewer component of Studio."""
import copy
from typing import Protocol
import dataclasses
from typing import Any
import mujoco
from mujoco.experimental.studio import endpoints
from mujoco.experimental.studio import handler_registry
from mujoco.experimental.studio import messages
from mujoco.experimental.studio import parser
from mujoco.experimental.studio import sim
@@ -28,25 +30,25 @@ import numpy as np
from mujoco.experimental.dear_imgui import dear_imgui as imgui
# ------------------------------------------------------------------------------
# Viewer application and custom message handlers
# ------------------------------------------------------------------------------
@dataclasses.dataclass(frozen=True)
class ViewerAppInitEvent(messages.Event):
"""Lifecycle event dispatched once when the ViewerApp is initialised.
Handlers that need access to the ViewerApp should handle this event
and cache the reference.
"""
viewer_app: 'ViewerApp'
class ViewerEventHandler(Protocol):
"""Invoked once per event received from the simulation."""
def handle(self, event: messages.Event) -> bool:
"""Return True if event was consumed, False to allow further processing."""
...
@dataclasses.dataclass(frozen=True)
class BuildGuiEvent(messages.Event):
"""Lifecycle event dispatched on every frame on the viewer side to build ImGui elements."""
class ViewerSnapshotHandler(Protocol):
"""Invoked once per snapshot received from the simulation."""
def handle(self, snapshot: messages.Snapshot) -> bool:
"""Return True if snapshot was consumed, False to allow further processing."""
...
@dataclasses.dataclass(frozen=True)
class UpdateEvent(messages.Event):
"""Lifecycle event dispatched on every frame on the viewer side before building GUI."""
class ViewerApp:
@@ -57,10 +59,16 @@ class ViewerApp:
viewer: viewer_protocol.Viewer,
endpoint: endpoints.ViewerEndpoint,
*,
viewer_event_handler: ViewerEventHandler | None = None,
viewer_snapshot_handler: ViewerSnapshotHandler | None = None,
):
"""Initializes the Studio application."""
handlers: list[Any] | None = None,
) -> None:
"""Initializes the Studio application.
Args:
viewer: A viewer display surface conforming to the Viewer protocol.
endpoint: The viewer endpoint for communication with the sim side.
handlers: Optional list of handler instances for viewer-side processing,
which are classes with methods decorated with ``@handler``.
"""
self.viewer = viewer
self.model: mujoco.MjModel = mujoco.MjSpec().compile()
self.data: mujoco.MjData = mujoco.MjData(self.model)
@@ -68,8 +76,11 @@ class ViewerApp:
self._last_model_id: int | None = id(self.model)
self.model_path: str = ''
self.endpoint = endpoint
self.viewer_event_handler = viewer_event_handler
self.viewer_snapshot_handler = viewer_snapshot_handler
# Instantiate handlers from user handlers + framework defaults.
all_handlers: list[Any] = list(handlers or [])
all_handlers.append(self)
self._handlers = handler_registry.HandlerRegistry(all_handlers)
self.step_control_state = sim.StepControl() # ONLY for state!
self.ux_state = ux.UxState()
@@ -82,6 +93,9 @@ class ViewerApp:
# pass it around with the camera would be convenient.
self._cam_speed = 0.001
# Dispatch lifecycle event so handlers can cache the app reference.
self._handlers.dispatch(ViewerAppInitEvent(viewer_app=self))
def close(self) -> None:
"""Signals the sim to exit and closes the viewer endpoint, releasing resources."""
if not self.should_quit:
@@ -112,7 +126,7 @@ class ViewerApp:
messages.MjOptionSnapshot(opt=copy.deepcopy(self.model.opt))
)
def handle_keyboard_events(self):
def handle_keyboard_events(self) -> None:
"""Handles keyboard events."""
is_freecam_wasd = self.ux_state.camera_index == ux.FREE_CAMERA_IDX
@@ -227,25 +241,7 @@ class ViewerApp:
# Process incoming events from the simulation.
incoming_events = self.endpoint.get_sim_events()
for event in incoming_events:
if (
self.viewer_event_handler is not None
and self.viewer_event_handler.handle(event)
):
continue
if isinstance(event, messages.ModelEvent):
# A new model was sent (initial load or hot-reload).
# Deep-copy the incoming model to ensure the viewer operates on its own
# isolated copy of the C++ struct. This is required if the sim/viewer
# run in different processes, and means we don't need to lock when
# running in the same process on different threads.
self.model = copy.deepcopy(event.model)
self.data = mujoco.MjData(self.model)
# Confirm model object is distinct from the incoming event.
assert id(self.model) != id(event.model)
elif isinstance(event, messages.ExitEvent):
self.should_quit = True
self.viewer.close()
self._handlers.dispatch(event)
# Detect model change from drop_file or ModelEvent (or external swap).
model_changed = False
@@ -262,18 +258,8 @@ class ViewerApp:
# Process incoming snapshots from the simulation process.
incoming_snapshots = self.endpoint.get_sim_snapshots()
for snapshot in incoming_snapshots:
if (
self.viewer_snapshot_handler is not None
and self.viewer_snapshot_handler.handle(snapshot)
):
continue
if not model_changed and isinstance(snapshot, messages.StateSnapshot):
state_size = mujoco.mj_stateSize(self.model, snapshot.state_sig)
if len(snapshot.state) == state_size:
mujoco.mj_setState(
self.model, self.data, snapshot.state, snapshot.state_sig
)
mujoco.mj_forward(self.model, self.data)
if not model_changed:
self._handlers.dispatch(snapshot)
self.handle_mouse_events()
self.handle_keyboard_events()
@@ -432,40 +418,43 @@ class ViewerApp:
imgui.End()
imgui.PopStyleVar(3)
@messages.handler(priority=messages.Priority.INTERNAL)
def _on_model(self, event: messages.ModelEvent) -> bool:
self.model = copy.deepcopy(event.model)
self.data = mujoco.MjData(self.model)
assert id(self.model) != id(event.model)
return True
# ------------------------------------------------------------------------------
# Viewer application run loop and gui/update customization hooks.
# ------------------------------------------------------------------------------
@messages.handler(priority=messages.Priority.INTERNAL)
def _on_exit(self, _: messages.ExitEvent) -> bool:
self.should_quit = True
self.viewer.close()
return True
@messages.handler(priority=messages.Priority.INTERNAL)
def _on_state(self, event: messages.StateSnapshot) -> bool:
state_size = mujoco.mj_stateSize(self.model, event.state_sig)
if len(event.state) == state_size:
mujoco.mj_setState(
self.model, self.data, event.state, event.state_sig
)
mujoco.mj_forward(self.model, self.data)
return True
class ViewerUpdateHook(Protocol):
"""Invoked once per viewer frame, after default updates ."""
@messages.handler(priority=messages.Priority.INTERNAL)
def _on_update(self, _: UpdateEvent) -> None:
self.update()
def update(self, app: 'ViewerApp') -> None:
"""Update custom state."""
...
class ViewerGuiHook(Protocol):
"""Invoked once per viewer frame after all update() calls.
By deferring GUI hooks until after all updates we ensure custom UI reflects
the latest application state.
"""
def build_gui(self, app: 'ViewerApp') -> None:
"""Draw custom UI."""
...
@messages.handler(priority=messages.Priority.INTERNAL)
def _on_build_gui(self, _: BuildGuiEvent) -> None:
self.build_gui()
def run_viewer(
viewer: viewer_protocol.Viewer,
viewer_endpoint: endpoints.ViewerEndpoint,
*,
viewer_gui_hook: ViewerGuiHook | None = None,
viewer_update_hook: ViewerUpdateHook | None = None,
viewer_event_handler: ViewerEventHandler | None = None,
viewer_snapshot_handler: ViewerSnapshotHandler | None = None,
handlers: list[Any] | None = None,
) -> None:
"""Run the viewer loop with the given viewer and endpoint.
@@ -479,36 +468,16 @@ def run_viewer(
Args:
viewer: A viewer display surface conforming to the Viewer protocol.
viewer_endpoint: The viewer endpoint for communication with the sim side.
viewer_gui_hook: Optional hook to draw custom ImGui panels.
viewer_update_hook: Optional hook called once per frame before GUI.
viewer_event_handler: Optional handler for simulation-to-viewer events.
viewer_snapshot_handler: Optional handler for sim-to-viewer snapshots.
handlers: Optional list of handler instances for viewer processing, which
are classes with methods decorated with ``@handler``.
"""
app = ViewerApp(
viewer,
viewer_endpoint,
viewer_event_handler=viewer_event_handler,
viewer_snapshot_handler=viewer_snapshot_handler,
)
# pylint: disable=protected-access
app = ViewerApp(viewer, viewer_endpoint, handlers=handlers)
# Viewer main loop.
while app.is_running():
# Update the viewer state and handle user input.
app.update()
# Update the viewer state.
if viewer_update_hook is not None:
viewer_update_hook.update(app)
# Draw the default Studio GUI.
app.build_gui()
# Allow custom GUI elements to be added.
if viewer_gui_hook is not None:
viewer_gui_hook.build_gui(app)
# Sync the viewer display surface with the current model and data.
app._handlers.dispatch(UpdateEvent())
app._handlers.dispatch(BuildGuiEvent())
app.viewer.sync(app.model, app.data)
app.close()