Don't pollute the global namespace when using mjpython.

Fixes #2265

PiperOrigin-RevId: 713442833
Change-Id: If155d9f57d4deef8b848478ebd8077a00417d5c5
This commit is contained in:
Saran Tunyasuvunakool
2025-01-08 15:47:15 -08:00
committed by Copybara-Service
parent 357ea024c0
commit 40ef08c8ed
2 changed files with 136 additions and 108 deletions
+135 -108
View File
@@ -86,95 +86,113 @@ void* mjpython_pymain(void* vargs) {
// Set up the condition variable to pass control back to the macOS main thread.
gil = cpython.PyGILState_Ensure();
cpython.PyRun_SimpleStringFlags("import threading; cond = threading.Condition()", nullptr);
cpython.PyRun_SimpleStringFlags(R"(
def _mjpython_make_cond():
# Don't pollute the global namespace.
global _mjpython_make_cond
del _mjpython_make_cond
import threading
global cond
cond = threading.Condition()
_mjpython_make_cond()
)", nullptr);
py_initialized.store(true);
// Wait until GLFW is initialized on macOS main thread, set up the queue and an atexit hook
// to enqueue a termination flag upon exit.
cpython.PyRun_SimpleStringFlags(R"(
import atexit
def _mjpython_init():
# Don't pollute the global namespace.
global _mjpython_init
del _mjpython_init
# The mujoco.viewer module should only be imported here after glfw.init() in the macOS main thread.
with cond:
cond.wait()
import mujoco.viewer
import atexit
import threading
# Similar to a queue.Queue(maxsize=1), but where only one active task is allowed at a time.
# With queue.Queue(1), another item is allowed to be enqueued before task_done is called.
class _MjPythonImpl(mujoco.viewer._MjPythonBase):
# The mujoco.viewer module should only be imported after glfw.init() in the macOS main thread.
with cond:
cond.wait()
import mujoco.viewer
# Termination statuses
NOT_TERMINATED = 0
TERMINATION_REQUESTED = 1
TERMINATION_ACCEPTED = 2
TERMINATED = 3
# Similar to a queue.Queue(maxsize=1), but where only one active task is allowed at a time.
# With queue.Queue(1), another item is allowed to be enqueued before task_done is called.
class _MjPythonImpl(mujoco.viewer._MjPythonBase):
def __init__(self):
self._cond = threading.Condition()
self._task = None
self._termination = self.__class__.NOT_TERMINATED
self._busy = False
# Termination statuses
NOT_TERMINATED = 0
TERMINATION_REQUESTED = 1
TERMINATION_ACCEPTED = 2
TERMINATED = 3
def launch_on_ui_thread(
self,
model,
data,
handle_return,
key_callback,
show_left_ui,
show_right_ui,
):
with self._cond:
if self._busy or self._task is not None:
raise RuntimeError('another MuJoCo viewer is already open')
else:
self._task = (
model,
data,
handle_return,
key_callback,
show_left_ui,
show_right_ui,
)
def __init__(self):
self._cond = threading.Condition()
self._task = None
self._termination = self.__class__.NOT_TERMINATED
self._busy = False
def launch_on_ui_thread(
self,
model,
data,
handle_return,
key_callback,
show_left_ui,
show_right_ui,
):
with self._cond:
if self._busy or self._task is not None:
raise RuntimeError('another MuJoCo viewer is already open')
else:
self._task = (
model,
data,
handle_return,
key_callback,
show_left_ui,
show_right_ui,
)
self._cond.notify()
def terminate(self):
with self._cond:
self._termination = self.__class__.TERMINATION_REQUESTED
self._cond.notify()
self._cond.wait_for(
lambda: self._termination == self.__class__.TERMINATED)
def get(self):
with self._cond:
self._cond.wait_for(
lambda: self._task is not None or self._termination)
if self._termination:
if self._termination == self.__class__.TERMINATION_REQUESTED:
self._termination = self.__class__.TERMINATION_ACCEPTED
return None
task = self._task
self._busy = True
self._task = None
return task
def done(self):
with self._cond:
self._busy = False
if self._termination == self.__class__.TERMINATION_ACCEPTED:
self._termination = self.__class__.TERMINATED
self._cond.notify()
def terminate(self):
with self._cond:
self._termination = self.__class__.TERMINATION_REQUESTED
self._cond.notify()
self._cond.wait_for(
lambda: self._termination == self.__class__.TERMINATED)
def get(self):
with self._cond:
self._cond.wait_for(
lambda: self._task is not None or self._termination)
mujoco.viewer._MJPYTHON = _MjPythonImpl()
atexit.register(mujoco.viewer._MJPYTHON.terminate)
if self._termination:
if self._termination == self.__class__.TERMINATION_REQUESTED:
self._termination = self.__class__.TERMINATION_ACCEPTED
return None
with cond:
cond.notify()
task = self._task
self._busy = True
self._task = None
return task
def done(self):
with self._cond:
self._busy = False
if self._termination == self.__class__.TERMINATION_ACCEPTED:
self._termination = self.__class__.TERMINATED
self._cond.notify()
mujoco.viewer._MJPYTHON = _MjPythonImpl()
atexit.register(mujoco.viewer._MJPYTHON.terminate)
del _MjPythonImpl # Don't pollute globals for user script.
with cond:
cond.notify()
del cond # Don't pollute globals for user script.
_mjpython_init()
)", nullptr);
// Run the Python interpreter main loop.
@@ -283,47 +301,56 @@ int main(int argc, char** argv) {
// to finish setting up _MJPYTHON, then serve incoming viewer launch requests.
PyGILState_STATE gil = cpython.PyGILState_Ensure();
cpython.PyRun_SimpleStringFlags(R"(
import ctypes
def _mjpython_main():
# Don't pollute the global namespace.
global _mjpython_main
del _mjpython_main
# GLFW must be initialized on the OS main thread (i.e. here).
import glfw
import mujoco.viewer
import ctypes
glfw.init()
glfw.poll_events()
ctypes.CDLL(None).mjpython_hide_dock_icon()
# GLFW must be initialized on the OS main thread (i.e. here).
import glfw
import mujoco.viewer
# Wait for Python main thread to finish setting up _MJPYTHON
with cond:
cond.notify()
cond.wait()
glfw.init()
glfw.poll_events()
ctypes.CDLL(None).mjpython_hide_dock_icon()
while True:
try:
# Wait for an incoming payload.
task = mujoco.viewer._MJPYTHON.get()
# Wait for Python main thread to finish setting up _MJPYTHON
global cond
with cond:
cond.notify()
cond.wait()
del cond
# None means that we are exiting.
if task is None:
glfw.terminate()
break
while True:
try:
# Wait for an incoming payload.
task = mujoco.viewer._MJPYTHON.get()
# Otherwise, launch the viewer.
model, data, handle_return, key_callback, show_left_ui, show_right_ui = task
ctypes.CDLL(None).mjpython_show_dock_icon()
mujoco.viewer._launch_internal(
model,
data,
run_physics_thread=False,
handle_return=handle_return,
key_callback=key_callback,
show_left_ui=show_left_ui,
show_right_ui=show_right_ui,
)
ctypes.CDLL(None).mjpython_hide_dock_icon()
# None means that we are exiting.
if task is None:
glfw.terminate()
break
finally:
mujoco.viewer._MJPYTHON.done()
# Otherwise, launch the viewer.
model, data, handle_return, key_callback, show_left_ui, show_right_ui = task
ctypes.CDLL(None).mjpython_show_dock_icon()
mujoco.viewer._launch_internal(
model,
data,
run_physics_thread=False,
handle_return=handle_return,
key_callback=key_callback,
show_left_ui=show_left_ui,
show_right_ui=show_right_ui,
)
ctypes.CDLL(None).mjpython_hide_dock_icon()
finally:
mujoco.viewer._MJPYTHON.done()
_mjpython_main()
)", nullptr);
cpython.PyGILState_Release(gil);