From 1f9dca8bc4cfbfc23f68c6bda7cdb6abebfdeb98 Mon Sep 17 00:00:00 2001 From: Alessio Quaglino Date: Mon, 7 Oct 2024 09:36:36 -0700 Subject: [PATCH] Do not expose plugin_slot. PiperOrigin-RevId: 683215401 Change-Id: Ic64f7221ed83ca3b0c8cca16eb680ea67c67b564 --- doc/APIreference/functions.rst | 2 +- doc/includes/references.h | 5 ++- include/mujoco/mjspec.h | 3 +- include/mujoco/mujoco.h | 4 +-- introspect/functions.py | 4 +-- introspect/structs.py | 7 +---- .../mujoco/codegen/generate_spec_bindings.py | 5 --- python/mujoco/specs_test.py | 3 -- src/user/user_api.cc | 4 +-- src/user/user_api.h | 4 +-- src/user/user_init.c | 1 - src/user/user_mesh.cc | 2 +- src/user/user_model.cc | 31 ++++++++++++++----- src/user/user_objects.cc | 14 ++++++--- src/user/user_objects.h | 3 +- src/xml/xml_native_reader.cc | 3 +- src/xml/xml_native_writer.cc | 4 +-- 17 files changed, 50 insertions(+), 49 deletions(-) diff --git a/doc/APIreference/functions.rst b/doc/APIreference/functions.rst index 0d58d0f9..43d4f84b 100644 --- a/doc/APIreference/functions.rst +++ b/doc/APIreference/functions.rst @@ -1516,7 +1516,7 @@ mjs_activatePlugin .. mujoco-include:: mjs_activatePlugin -Activate plugin, return slot number. +Activate plugin. .. _Errorandmemory: diff --git a/doc/includes/references.h b/doc/includes/references.h index a79f314f..de2d533b 100644 --- a/doc/includes/references.h +++ b/doc/includes/references.h @@ -1739,9 +1739,8 @@ typedef struct mjsOrientation_ { // alternative orientation specifiers } mjsOrientation; typedef struct mjsPlugin_ { // plugin specification mjsElement* element; // element type - mjString* name; // name + mjString* name; // instance name mjString* plugin_name; // plugin name - int plugin_slot; // global registered slot number of the plugin mjtByte active; // is the plugin active mjString* info; // message appended to compiler errors } mjsPlugin; @@ -3171,7 +3170,7 @@ int mj_setLengthRange(mjModel* m, mjData* d, int index, mjSpec* mj_makeSpec(void); mjSpec* mj_copySpec(const mjSpec* s); void mj_deleteSpec(mjSpec* s); -int mjs_activatePlugin(mjSpec* s, const char* name); +void mjs_activatePlugin(mjSpec* s, const char* name); void mj_printFormattedModel(const mjModel* m, const char* filename, const char* float_format); void mj_printModel(const mjModel* m, const char* filename); void mj_printFormattedData(const mjModel* m, mjData* d, const char* filename, diff --git a/include/mujoco/mjspec.h b/include/mujoco/mjspec.h index 4c706338..c4e9cafd 100644 --- a/include/mujoco/mjspec.h +++ b/include/mujoco/mjspec.h @@ -182,9 +182,8 @@ typedef struct mjsOrientation_ { // alternative orientation specifiers typedef struct mjsPlugin_ { // plugin specification mjsElement* element; // element type - mjString* name; // name + mjString* name; // instance name mjString* plugin_name; // plugin name - int plugin_slot; // global registered slot number of the plugin mjtByte active; // is the plugin active mjString* info; // message appended to compiler errors } mjsPlugin; diff --git a/include/mujoco/mujoco.h b/include/mujoco/mujoco.h index 771209f1..879ae85a 100644 --- a/include/mujoco/mujoco.h +++ b/include/mujoco/mujoco.h @@ -241,8 +241,8 @@ MJAPI mjSpec* mj_copySpec(const mjSpec* s); // Free memory allocation in mjSpec. MJAPI void mj_deleteSpec(mjSpec* s); -// Activate plugin, return slot number. -MJAPI int mjs_activatePlugin(mjSpec* s, const char* name); +// Activate plugin. +MJAPI void mjs_activatePlugin(mjSpec* s, const char* name); //---------------------------------- Printing ------------------------------------------------------ diff --git a/introspect/functions.py b/introspect/functions.py index 7f7d6133..0d576b8e 100644 --- a/introspect/functions.py +++ b/introspect/functions.py @@ -1048,7 +1048,7 @@ FUNCTIONS: Mapping[str, FunctionDecl] = dict([ ('mjs_activatePlugin', FunctionDecl( name='mjs_activatePlugin', - return_type=ValueType(name='int'), + return_type=ValueType(name='void'), parameters=( FunctionParameterDecl( name='s', @@ -1063,7 +1063,7 @@ FUNCTIONS: Mapping[str, FunctionDecl] = dict([ ), ), ), - doc='Activate plugin, return slot number.', + doc='Activate plugin.', )), ('mj_printFormattedModel', FunctionDecl( diff --git a/introspect/structs.py b/introspect/structs.py index 2b79c85f..9455d2a7 100644 --- a/introspect/structs.py +++ b/introspect/structs.py @@ -8576,7 +8576,7 @@ STRUCTS: Mapping[str, StructDecl] = dict([ type=PointerType( inner_type=ValueType(name='mjString'), ), - doc='name', + doc='instance name', ), StructFieldDecl( name='plugin_name', @@ -8585,11 +8585,6 @@ STRUCTS: Mapping[str, StructDecl] = dict([ ), doc='plugin name', ), - StructFieldDecl( - name='plugin_slot', - type=ValueType(name='int'), - doc='global registered slot number of the plugin', - ), StructFieldDecl( name='active', type=ValueType(name='mjtByte'), diff --git a/python/mujoco/codegen/generate_spec_bindings.py b/python/mujoco/codegen/generate_spec_bindings.py index 421e5d77..e02733c1 100644 --- a/python/mujoco/codegen/generate_spec_bindings.py +++ b/python/mujoco/codegen/generate_spec_bindings.py @@ -461,11 +461,6 @@ def generate_add() -> None: } catch (const py::cast_error &e) { throw pybind11::value_error("plugin.instance_name should be a string."); } - try { - plugin.plugin_slot = input->plugin_slot; - } catch (const py::cast_error &e) { - throw pybind11::value_error("plugin.plugin_slot should be an int."); - } try { plugin.active = input->active; } catch (const py::cast_error &e) { diff --git a/python/mujoco/specs_test.py b/python/mujoco/specs_test.py index 26938c8c..acce44d1 100644 --- a/python/mujoco/specs_test.py +++ b/python/mujoco/specs_test.py @@ -207,13 +207,11 @@ class SpecsTest(absltest.TestCase): plugin = spec.add_plugin( name='instance_name', plugin_name='mujoco.plugin', - plugin_slot=7, active=True, info='info', ) self.assertEqual(plugin.name, 'instance_name') self.assertEqual(plugin.plugin_name, 'mujoco.plugin') - self.assertEqual(plugin.plugin_slot, 7) self.assertEqual(plugin.active, True) self.assertEqual(plugin.info, 'info') @@ -229,7 +227,6 @@ class SpecsTest(absltest.TestCase): body_with_plugin = spec.worldbody.add_body(plugin=plugin) self.assertEqual(body_with_plugin.plugin.name, 'instance_name') self.assertEqual(body_with_plugin.plugin.plugin_name, 'mujoco.plugin') - self.assertEqual(body_with_plugin.plugin.plugin_slot, 7) self.assertEqual(body_with_plugin.plugin.active, True) self.assertEqual(body_with_plugin.plugin.info, 'info') diff --git a/src/user/user_api.cc b/src/user/user_api.cc index d6ada1d9..9dba3c81 100644 --- a/src/user/user_api.cc +++ b/src/user/user_api.cc @@ -206,16 +206,14 @@ void mjs_addSpec(mjSpec* s, mjSpec* child) { // activate plugin -int mjs_activatePlugin(mjSpec* s, const char* name) { +void mjs_activatePlugin(mjSpec* s, const char* name) { int plugin_slot = -1; const mjpPlugin* plugin = mjp_getPlugin(name, &plugin_slot); if (!plugin) { mju_error("unknown plugin '%s'", name); - return -1; } mjCModel* model = static_cast(s->element); model->ActivatePlugin(plugin, plugin_slot); - return plugin_slot; } diff --git a/src/user/user_api.h b/src/user/user_api.h index 162097b4..f3f25ebc 100644 --- a/src/user/user_api.h +++ b/src/user/user_api.h @@ -63,8 +63,8 @@ MJAPI void mj_deleteSpec(mjSpec* s); // Add spec (model asset) to spec. MJAPI void mjs_addSpec(mjSpec* s, mjSpec* child); -// Activate plugin, return slot number. -MJAPI int mjs_activatePlugin(mjSpec* s, const char* name); +// Activate plugin. +MJAPI void mjs_activatePlugin(mjSpec* s, const char* name); //---------------------------------- Attachment ---------------------------------------------------- diff --git a/src/user/user_init.c b/src/user/user_init.c index e031831a..59ae0cc0 100644 --- a/src/user/user_init.c +++ b/src/user/user_init.c @@ -402,6 +402,5 @@ void mjs_defaultKey(mjsKey* key) { // default plugin attributes void mjs_defaultPlugin(mjsPlugin* plugin) { memset(plugin, 0, sizeof(mjsPlugin)); - plugin->plugin_slot = -1; } diff --git a/src/user/user_mesh.cc b/src/user/user_mesh.cc index 91c0a99c..5b7174d4 100644 --- a/src/user/user_mesh.cc +++ b/src/user/user_mesh.cc @@ -294,7 +294,7 @@ void mjCMesh::LoadSDF() { mjCPlugin* plugin_instance = static_cast(plugin.element); model->ResolvePlugin(this, plugin_name, plugin_instance_name, &plugin_instance); plugin.element = plugin_instance; - const mjpPlugin* pplugin = mjp_getPluginAtSlot(plugin_instance->spec.plugin_slot); + const mjpPlugin* pplugin = mjp_getPluginAtSlot(plugin_instance->plugin_slot); if (!(pplugin->capabilityflags & mjPLUGIN_SDF)) { throw mjCError(this, "plugin '%s' does not support signed distance fields", pplugin->name); } diff --git a/src/user/user_model.cc b/src/user/user_model.cc index 58b1e605..d8b4256c 100644 --- a/src/user/user_model.cc +++ b/src/user/user_model.cc @@ -315,6 +315,7 @@ void mjCModel::CopyPlugin(std::vector& dest, continue; } mjCPlugin* candidate = new mjCPlugin(*plugin); + candidate->model = this; candidate->NameSpace(plugin->model); bool referenced = instances.find(candidate->name) != instances.end(); auto same_name = [candidate](const mjCPlugin* dest) { return dest->name == candidate->name; }; @@ -3973,7 +3974,7 @@ void mjCModel::TryCompile(mjModel*& m, mjData*& d, const mjVFS* vfs) { { int adr = 0; for (int i = 0; i < nplugin; ++i) { - m->plugin[i] = plugins_[i]->spec.plugin_slot; + m->plugin[i] = plugins_[i]->plugin_slot; const int size = plugins_[i]->flattened_attributes.size(); std::memcpy(m->plugin_attr + adr, plugins_[i]->flattened_attributes.data(), size); @@ -4499,24 +4500,37 @@ void mjCModel::ActivatePlugin(const mjpPlugin* plugin, int slot) { void mjCModel::ResolvePlugin(mjCBase* obj, const std::string& plugin_name, const std::string& plugin_instance_name, mjCPlugin** plugin_instance) { + std::string pname = plugin_name; + + // if the plugin name is not specified by the user, infer it from the plugin instance + if (plugin_name.empty() && !plugin_instance_name.empty()) { + mjCBase* plugin_obj = FindObject(mjOBJ_PLUGIN, plugin_instance_name); + if (plugin_obj) { + pname = static_cast(plugin_obj)->plugin_name; + } else { + throw mjCError(obj, "unrecognized name '%s' for plugin instance", + plugin_instance_name.c_str()); + } + } + // if plugin_name is specified, check if it is in the list of active plugins // (in XML, active plugins are those declared as ) int plugin_slot = -1; - if (!plugin_name.empty()) { + if (!pname.empty()) { for (int i = 0; i < active_plugins_.size(); ++i) { - if (active_plugins_[i].first->name == plugin_name) { + if (active_plugins_[i].first->name == pname) { plugin_slot = active_plugins_[i].second; break; } } if (plugin_slot == -1) { - throw mjCError(obj, "unrecognized plugin '%s'", plugin_name.c_str()); + throw mjCError(obj, "unrecognized plugin '%s'", pname.c_str()); } } // implicit plugin instance - if (*plugin_instance && (*plugin_instance)->spec.plugin_slot == -1) { - (*plugin_instance)->spec.plugin_slot = plugin_slot; + if (*plugin_instance && (*plugin_instance)->plugin_slot == -1) { + (*plugin_instance)->plugin_slot = plugin_slot; (*plugin_instance)->parent = obj; } @@ -4524,14 +4538,15 @@ void mjCModel::ResolvePlugin(mjCBase* obj, const std::string& plugin_name, else if (!*plugin_instance) { *plugin_instance = static_cast(FindObject(mjOBJ_PLUGIN, plugin_instance_name)); + (*plugin_instance)->plugin_slot = plugin_slot; if (!*plugin_instance) { throw mjCError( obj, "unrecognized name '%s' for plugin instance", plugin_instance_name.c_str()); } - if (plugin_slot != -1 && plugin_slot != (*plugin_instance)->spec.plugin_slot) { + if (plugin_slot != -1 && plugin_slot != (*plugin_instance)->plugin_slot) { throw mjCError( obj, "'plugin' attribute does not match that of the instance"); } - plugin_slot = (*plugin_instance)->spec.plugin_slot; + plugin_slot = (*plugin_instance)->plugin_slot; } } diff --git a/src/user/user_objects.cc b/src/user/user_objects.cc index 8218c6a4..bf65075b 100644 --- a/src/user/user_objects.cc +++ b/src/user/user_objects.cc @@ -1695,7 +1695,7 @@ void mjCBody::Compile(void) { mjCPlugin* plugin_instance = static_cast(plugin.element); model->ResolvePlugin(this, plugin_name, plugin_instance_name, &plugin_instance); plugin.element = plugin_instance; - const mjpPlugin* pplugin = mjp_getPluginAtSlot(plugin_instance->spec.plugin_slot); + const mjpPlugin* pplugin = mjp_getPluginAtSlot(plugin_instance->plugin_slot); if (!(pplugin->capabilityflags & mjPLUGIN_PASSIVE)) { throw mjCError(this, "plugin '%s' does not support passive forces", pplugin->name); } @@ -2949,7 +2949,7 @@ void mjCGeom::Compile(void) { mjCPlugin* plugin_instance = static_cast(plugin.element); model->ResolvePlugin(this, plugin_name, plugin_instance_name, &plugin_instance); plugin.element = plugin_instance; - const mjpPlugin* pplugin = mjp_getPluginAtSlot(plugin_instance->spec.plugin_slot); + const mjpPlugin* pplugin = mjp_getPluginAtSlot(plugin_instance->plugin_slot); if (!(pplugin->capabilityflags & mjPLUGIN_SDF)) { throw mjCError(this, "plugin '%s' does not support sign distance fields", pplugin->name); } @@ -5908,7 +5908,7 @@ void mjCActuator::Compile(void) { mjCPlugin* plugin_instance = static_cast(plugin.element); model->ResolvePlugin(this, plugin_name, plugin_instance_name, &plugin_instance); plugin.element = plugin_instance; - const mjpPlugin* pplugin = mjp_getPluginAtSlot(plugin_instance->spec.plugin_slot); + const mjpPlugin* pplugin = mjp_getPluginAtSlot(plugin_instance->plugin_slot); if (!(pplugin->capabilityflags & mjPLUGIN_ACTUATOR)) { throw mjCError(this, "plugin '%s' does not support actuators", pplugin->name); } @@ -6415,7 +6415,7 @@ void mjCSensor::Compile(void) { mjCPlugin* plugin_instance = static_cast(plugin.element); model->ResolvePlugin(this, plugin_name, plugin_instance_name, &plugin_instance); plugin.element = plugin_instance; - const mjpPlugin* pplugin = mjp_getPluginAtSlot(plugin_instance->spec.plugin_slot); + const mjpPlugin* pplugin = mjp_getPluginAtSlot(plugin_instance->plugin_slot); if (!(pplugin->capabilityflags & mjPLUGIN_SENSOR)) { throw mjCError(this, "plugin '%s' does not support sensors", pplugin->name); } @@ -6918,6 +6918,7 @@ void mjCKey::Compile(const mjModel* m) { mjCPlugin::mjCPlugin(mjCModel* _model) { name = ""; nstate = -1; + plugin_slot = -1; parent = this; model = _model; name.clear(); @@ -6945,6 +6946,7 @@ mjCPlugin& mjCPlugin::operator=(const mjCPlugin& other) { this->spec = other.spec; *static_cast(this) = static_cast(other); parent = this; + plugin_slot = other.plugin_slot; } return *this; } @@ -6953,7 +6955,9 @@ mjCPlugin& mjCPlugin::operator=(const mjCPlugin& other) { // compiler void mjCPlugin::Compile(void) { - const mjpPlugin* plugin = mjp_getPluginAtSlot(spec.plugin_slot); + mjCPlugin* plugin_instance = this; + model->ResolvePlugin(this, plugin_name, name, &plugin_instance); + const mjpPlugin* plugin = mjp_getPluginAtSlot(plugin_slot); // clear precompiled flattened_attributes.clear(); diff --git a/src/user/user_objects.h b/src/user/user_objects.h index ad29d1ae..82641520 100644 --- a/src/user/user_objects.h +++ b/src/user/user_objects.h @@ -1424,7 +1424,8 @@ class mjCPlugin : public mjCPlugin_ { mjCPlugin(const mjCPlugin& other); mjCPlugin& operator=(const mjCPlugin& other); mjsPlugin spec; - mjCBase* parent; // parent object (only used when generating error message) + mjCBase* parent; // parent object (only used when generating error message) + int plugin_slot; // global registered slot number of the plugin private: void Compile(void); // compiler diff --git a/src/xml/xml_native_reader.cc b/src/xml/xml_native_reader.cc index 53f4cb83..4bcd7cd1 100644 --- a/src/xml/xml_native_reader.cc +++ b/src/xml/xml_native_reader.cc @@ -2890,7 +2890,7 @@ void mjXReader::Extension(XMLElement* section) { if (name == "plugin") { string plugin_name; ReadAttrTxt(elem, "plugin", plugin_name, /* required = */ true); - int plugin_slot = mjs_activatePlugin(spec, plugin_name.c_str()); + mjs_activatePlugin(spec, plugin_name.c_str()); XMLElement* child = FirstChildElement(elem); while (child) { @@ -2909,7 +2909,6 @@ void mjXReader::Extension(XMLElement* section) { throw mjXError(child, "plugin instance must have a name"); } ReadPluginConfigs(child, p); - p->plugin_slot = plugin_slot; } child = NextSiblingElement(child); } diff --git a/src/xml/xml_native_writer.cc b/src/xml/xml_native_writer.cc index d0aea40a..54c18d89 100644 --- a/src/xml/xml_native_writer.cc +++ b/src/xml/xml_native_writer.cc @@ -827,7 +827,7 @@ void mjXWriter::OnePlugin(XMLElement* elem, const mjsPlugin* plugin) { } else { WriteAttrTxt(elem, "plugin", plugin_name); const mjpPlugin* pplugin = mjp_getPluginAtSlot( - static_cast(plugin->element)->spec.plugin_slot); + static_cast(plugin->element)->plugin_slot); const char* c = &(static_cast(plugin->element)->flattened_attributes[0]); for (int i = 0; i < pplugin->nattribute; ++i) { string value(c); @@ -1338,7 +1338,7 @@ void mjXWriter::Extension(XMLElement* root) { } // check if we need to open a new section - const mjpPlugin* plugin = mjp_getPluginAtSlot(pp->spec.plugin_slot); + const mjpPlugin* plugin = mjp_getPluginAtSlot(pp->plugin_slot); if (plugin != last_plugin) { plugin_elem = InsertEnd(section, "plugin"); WriteAttrTxt(plugin_elem, "plugin", plugin->name);