Files
Mujoco_WASM/python/mujoco/experimental/studio/endpoints.py
T
Matija Kecman e0bb2b0665 Introduce message and endpoint abstractions for MuJoCo Studio communication.
-   `messages.py`: Defines base classes for `Message`, `Snapshot`, and `Event`, along with `SnapshotChannel` and `EventChannel` protocols. It also includes concrete message types for state, reset, model, perturbation, pause state, and exit events.
-   `endpoints.py`: Provides `ViewerEndpoint` and `SimEndpoint` classes to manage message routing between the simulation and the viewer via the defined channels. A `make_endpoints` function is included to create both endpoints.

PiperOrigin-RevId: 939659912
Change-Id: I61a52a7d25431bc4ed38a184813882145eb421d1
2026-06-29 00:24:15 -07:00

133 lines
4.1 KiB
Python

# Copyright 2026 DeepMind Technologies Limited
#
# 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
#
# https://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.
"""Endpoints route messages to the right channel."""
from mujoco.experimental.studio import messages as vp
class ViewerEndpoint:
"""Viewer endpoint to route Messages to the Sim using the right Channel.
The viewer has an outgoing EventChannel but no outgoing SnapshotChannel.
The viewer has an incoming EventChannel and SnapshotChannel.
"""
def __init__(
self,
*,
s2v_snapshot: vp.SnapshotChannel,
s2v_events: vp.EventChannel,
v2s_events: vp.EventChannel,
):
self._s2v_snapshot = s2v_snapshot
self._s2v_events = s2v_events
self._v2s_events = v2s_events
self._is_closed = False
def send_to_sim(self, message: vp.Message) -> None:
if self._is_closed:
raise RuntimeError('ViewerEndpoint is closed')
if isinstance(message, vp.Event):
self._v2s_events.put(message)
elif isinstance(message, vp.Snapshot):
raise TypeError('viewer cannot produce snapshots')
else:
raise TypeError(f'expected Event, got {type(message).__name__}')
def get_sim_events(self) -> list[vp.Event]:
"""Returns all pending events from the simulation."""
if self._is_closed:
raise RuntimeError('ViewerEndpoint is closed')
return self._s2v_events.get()
def get_sim_snapshots(self) -> list[vp.Snapshot]:
"""Returns all pending latest snapshots from the simulation, one per type."""
if self._is_closed:
raise RuntimeError('ViewerEndpoint is closed')
return self._s2v_snapshot.get()
def close(self) -> None:
"""Close all channels owned by this endpoint (idempotent)."""
if not self._is_closed:
self._is_closed = True
self._v2s_events.close()
self._s2v_events.close()
self._s2v_snapshot.close()
class SimEndpoint:
"""Sim endpoint to route Messages to the Viewer using the right Channel.
The simulation has an outgoing SnapshotChannel and EventChannel.
The simulation has an incoming EventChannel but no incoming SnapshotChannel.
"""
def __init__(
self,
*,
s2v_snapshot: vp.SnapshotChannel,
s2v_events: vp.EventChannel,
v2s_events: vp.EventChannel,
):
self._s2v_snapshot = s2v_snapshot
self._s2v_events = s2v_events
self._v2s_events = v2s_events
self._is_closed = False
def send_to_viewer(self, message: vp.Message) -> None:
if self._is_closed:
raise RuntimeError('SimEndpoint is closed')
if isinstance(message, vp.Event):
self._s2v_events.put(message)
elif isinstance(message, vp.Snapshot):
self._s2v_snapshot.put(message)
else:
raise TypeError(
f'expected Snapshot or Event, got {type(message).__name__}'
)
def get_viewer_events(self) -> list[vp.Event]:
"""Returns all pending events from the viewer."""
if self._is_closed:
raise RuntimeError('SimEndpoint is closed')
return self._v2s_events.get()
def close(self) -> None:
"""Close all channels owned by this endpoint."""
if not self._is_closed:
self._is_closed = True
self._s2v_snapshot.close()
self._s2v_events.close()
self._v2s_events.close()
def make_endpoints(
*,
s2v_snapshot: vp.SnapshotChannel,
s2v_events: vp.EventChannel,
v2s_events: vp.EventChannel,
) -> tuple[ViewerEndpoint, SimEndpoint]:
"""Returns viewer and simulation endpoints."""
viewer_endpoint = ViewerEndpoint(
s2v_snapshot=s2v_snapshot,
s2v_events=s2v_events,
v2s_events=v2s_events,
)
sim_endpoint = SimEndpoint(
s2v_snapshot=s2v_snapshot,
s2v_events=s2v_events,
v2s_events=v2s_events,
)
return viewer_endpoint, sim_endpoint