Make MuJoCo Python bindings compatible with free-threading.

Introduce a new header `gil.h` defining `MutexLockIfGilDisabled` to support thread-safety in both standard and free-threaded CPython builds.

Protect critical shared states and registries:
- Guard global Python callback pointers in `callbacks.cc` using a mutex. Move `gil_scoped_acquire` into local blocks around refcount modifications to prevent `longjmp` from bypassing destructors.
- Protect raw pointer maps in `structs_wrappers.cc` with static mutexes.
- Replace TOCTOU race in `mjcb_time` initialization with thread-safe `std::call_once`.
- Add synchronization to lazy indexer array cache initialization in `indexers.cc` and `indexer_xmacro.h`.
- Protect vector mutations in `StructListBase::PopulateUpTo` in `structs.h` with a mutex.
- Revert unnecessary atomic changes to threadpool counters.
- Declare free-threading compatibility by passing `pybind11::mod_gil_not_used()` to all extension modules.

Fixes #3259
Fixes #3256
Fixes #2978

PiperOrigin-RevId: 941101502
Change-Id: Iec4ce58afcbc75d4b0be6a9a21fc8a47854242e3
This commit is contained in:
Saran Tunyasuvunakool
2026-07-01 08:08:10 -07:00
committed by Copybara-Service
parent cab191755a
commit a07ae6f849
21 changed files with 489 additions and 66 deletions
+7 -2
View File
@@ -19,6 +19,7 @@
#include <utility>
#include <vector>
#include "gil.h"
#include "indexers.h"
#include "raw.h"
#include "util/crossplatform.h"
@@ -184,7 +185,7 @@ MjModelIndexer::MjModelIndexer(raw::MjModel* m, py::handle owner)
name_to_id_(*m),
id_to_name_(*m)
#define XGROUP(MjModelFieldGroupedViews, field, nfield, FIELD_XMACROS) \
, field##_(m->nfield, std::nullopt)
, field##_(m->nfield)
MJMODEL_VIEW_GROUPS
#undef XGROUP
{}
@@ -194,6 +195,7 @@ MjModelIndexer::MjModelIndexer(raw::MjModel* m, py::handle owner)
if (i >= field##_.size() || i < 0) { \
throw py::index_error(IndexErrorMessage(i, field##_.size())); \
} \
MutexLockIfGilDisabled lock(lazy_init_mutex_); \
auto& indexer = field##_[i]; \
if (!indexer.has_value()) { \
const std::string& name = id_to_name_.field[i]; \
@@ -224,7 +226,7 @@ MjDataIndexer::MjDataIndexer(raw::MjData* d, const raw::MjModel* m,
name_to_id_(*m),
id_to_name_(*m)
#define XGROUP(MjDataGroupedViews, field, nfield, FIELD_XMACROS) \
, field##_(m->nfield, std::nullopt)
, field##_(m->nfield)
MJDATA_VIEW_GROUPS
#undef XGROUP
{}
@@ -234,6 +236,7 @@ MjDataIndexer::MjDataIndexer(raw::MjData* d, const raw::MjModel* m,
if (i >= field##_.size() || i < 0) { \
throw py::index_error(IndexErrorMessage(i, field##_.size())); \
} \
MutexLockIfGilDisabled lock(lazy_init_mutex_); \
auto& indexer = field##_[i]; \
if (!indexer.has_value()) { \
const std::string& name = id_to_name_.field[i]; \
@@ -270,6 +273,7 @@ MJDATA_VIEW_GROUPS
#define MJ_M(n) m_->n
#define X(type, prefix, var, dim0, dim1) \
py::array_t<type> XGROUP::var() { \
MutexLockIfGilDisabled lock(lazy_init_mutex_); \
if (!var##_.has_value()) { \
var##_.emplace(MakeArray<&raw::MjModel::dim0>( \
m_->prefix##var, index_, MAKE_SHAPE(dim1), *m_, owner_)); \
@@ -361,6 +365,7 @@ MJMODEL_KEYFRAME
#define MJ_M(n) m_->n
#define X(type, prefix, var, dim0, dim1) \
py::array_t<type> XGROUP::var() { \
MutexLockIfGilDisabled lock(lazy_init_mutex_); \
if (!var##_.has_value()) { \
var##_.emplace(MakeArray<&raw::MjModel::dim0>( \
d_->prefix##var, index_, MAKE_SHAPE(dim1), *m_, owner_)); \