From e6e6e3b51b719e361f9618c1635f1d3169825eaf Mon Sep 17 00:00:00 2001 From: Alessio Quaglino Date: Mon, 30 Sep 2024 08:06:54 -0700 Subject: [PATCH] Copy plugins when attaching a model. Before this CL, the same mjCPlugin objects were referenced by both models. Fixes #2100. PiperOrigin-RevId: 680579479 Change-Id: Ida3459f78f9fa4a942cea31b3d5d3e8da2832feb --- src/user/user_model.cc | 46 +++++++++++++++++++++++++---------- src/user/user_model.h | 1 - test/user/user_api_test.cc | 50 ++++++++++++++++++++++++++++++++++++++ 3 files changed, 83 insertions(+), 14 deletions(-) diff --git a/src/user/user_model.cc b/src/user/user_model.cc index ff9d6888..ec8ae6aa 100644 --- a/src/user/user_model.cc +++ b/src/user/user_model.cc @@ -24,11 +24,11 @@ #include #include #include -#include #include #include #include #include +#include #include #include @@ -164,9 +164,6 @@ mjCModel::mjCModel() { // point to model from spec PointToLocal(); - - // this class allocated the plugins - plugin_owner = true; } @@ -180,7 +177,6 @@ mjCModel::mjCModel(const mjCModel& other) { mjCModel& mjCModel::operator=(const mjCModel& other) { if (this != &other) { - plugin_owner = false; this->spec = other.spec; *static_cast(this) = static_cast(other); *static_cast(this) = static_cast(other); @@ -302,6 +298,22 @@ void mjCModel::SaveDofOffsets() { +template +static void mapplugin( + const std::unordered_map& plugin_map, std::vector& list) { + for (const auto& element : list) { + if (element->spec.plugin.element) { + mjCPlugin* plugin = static_cast(element->spec.plugin.element); + // the referenced plugin might already exist in the source model + if (plugin_map.find(plugin) != plugin_map.end()) { + element->spec.plugin.element = plugin_map.at(plugin); + } + } + } +} + + + mjCModel& mjCModel::operator+=(const mjCModel& other) { // create global lists mjCBody *world = bodies_[0]; @@ -323,6 +335,21 @@ mjCModel& mjCModel::operator+=(const mjCModel& other) { for (const auto& key : other.key_pending_) { key_pending_.push_back(key); } + + // create new plugins and map them + std::unordered_map plugin_map; + for (const auto& plugin : other.plugins_) { + plugins_.push_back(new mjCPlugin(*plugin)); + plugin_map[plugin] = plugins_.back(); + } + mapplugin(plugin_map, bodies_); + mapplugin(plugin_map, geoms_); + mapplugin(plugin_map, meshes_); + mapplugin(plugin_map, actuators_); + mapplugin(plugin_map, sensors_); + for (const auto& active_plugin : other.active_plugins_) { + active_plugins_.emplace_back(active_plugin); + } } CopyList(flexes_, other.flexes_); CopyList(pairs_, other.pairs_); @@ -335,10 +362,6 @@ mjCModel& mjCModel::operator+=(const mjCModel& other) { CopyList(texts_, other.texts_); CopyList(tuples_, other.tuples_); - // plugins are global - plugins_ = other.plugins_; - active_plugins_ = other.active_plugins_; - // restore to the original state if (!compiled) { ResetTreeLists(); @@ -613,10 +636,7 @@ mjCModel::~mjCModel() { for (int i=0; i key_pending_; // attached keyframes }; #endif // MUJOCO_SRC_USER_USER_MODEL_H_ diff --git a/test/user/user_api_test.cc b/test/user/user_api_test.cc index 1f2b309a..9cfc823d 100644 --- a/test/user/user_api_test.cc +++ b/test/user/user_api_test.cc @@ -216,6 +216,56 @@ TEST_F(PluginTest, DeletePlugin) { mj_deleteModel(newmodel); } +TEST_F(PluginTest, AttachPlugin) { + static constexpr char xml_1[] = R"( + + + + + )"; + + static constexpr char xml_2[] = R"( + + + + + + + + + + + + + + + + + + + )"; + + std::array err; + mjSpec* spec_1 = mj_parseXMLString(xml_1, 0, err.data(), err.size()); + ASSERT_THAT(spec_1, NotNull()) << err.data(); + mjSpec* spec_2 = mj_parseXMLString(xml_2, 0, err.data(), err.size()); + ASSERT_THAT(spec_2, NotNull()) << err.data(); + + mjsBody* body_1 = mjs_findBody(spec_1, "body"); + EXPECT_THAT(body_1, NotNull()); + mjsFrame* attachment_frame = mjs_addFrame(body_1, 0); + EXPECT_THAT(attachment_frame, NotNull()); + + mjs_attachBody(attachment_frame, mjs_findBody(spec_2, "body"), "child-", ""); + mjModel* model = mj_compile(spec_1, nullptr); + EXPECT_THAT(model, NotNull()); + + mj_deleteModel(model); + mj_deleteSpec(spec_1); + mj_deleteSpec(spec_2); +} + TEST_F(MujocoTest, RecompileFails) { mjSpec* spec = mj_makeSpec(); mjsBody* body = mjs_addBody(mjs_findBody(spec, "world"), 0);