diff --git a/python/mujoco/codegen/generate_spec_bindings.py b/python/mujoco/codegen/generate_spec_bindings.py index 6ec1e24c..4dd24eea 100644 --- a/python/mujoco/codegen/generate_spec_bindings.py +++ b/python/mujoco/codegen/generate_spec_bindings.py @@ -93,16 +93,26 @@ def _value_binding_code( fulltype = fulltype.replace('mjOption', 'raw::MjOption') fulltype = fulltype.replace('mjVisual', 'raw::MjVisual') fulltype = fulltype.replace('mjStatistic', 'raw::MjStatistic') - element = '.element' if fullvarname == 'plugin' else '' + element = '' + + if field.name == 'mjsPlugin': + setter = f"""[]({rawclassname}& self, {fulltype} {varname}) {{ + if (self.{fullvarname}.name && {varname}.name) *self.{fullvarname}.name = *{varname}.name; + if (self.{fullvarname}.plugin_name && {varname}.plugin_name) *self.{fullvarname}.plugin_name = *{varname}.plugin_name; + self.{fullvarname}.active = {varname}.active; + if (self.{fullvarname}.info && {varname}.info) *self.{fullvarname}.info = *{varname}.info; + }}""" + else: + setter = f"""[]({rawclassname}& self, {fulltype} {varname}) {{ + self.{fullvarname}{element} = {varname}{element}; + }}""" def_property_args = ( f'"{varname}"', f"""[]({rawclassname}& self) -> {fulltype} {{ return self.{fullvarname}; }}""", - f"""[]({rawclassname}& self, {fulltype} {varname}) {{ - self.{fullvarname}{element} = {varname}{element}; - }}""", + setter, ) if field.name not in SCALAR_TYPES: diff --git a/python/mujoco/specs_test.py b/python/mujoco/specs_test.py index d77308fc..0d25c838 100644 --- a/python/mujoco/specs_test.py +++ b/python/mujoco/specs_test.py @@ -996,6 +996,23 @@ class SpecsTest(absltest.TestCase): self.assertIsNotNone(model) self.assertEqual(model.nplugin, 0) + def testPluginAssignment(self): + spec = mujoco.MjSpec() + body = spec.worldbody.add_body() + body.name = 'test_body' + + plugin = spec.add_plugin( + name='test_instance', plugin_name='mujoco.elasticity.cable', active=True + ) + plugin.config = {'twist': '1e2', 'bend': '4e1'} + + # Assignment should copy active flag and names + body.plugin = plugin + + self.assertTrue(body.plugin.active) + self.assertEqual(body.plugin.name, 'test_instance') + self.assertEqual(body.plugin.plugin_name, 'mujoco.elasticity.cable') + def test_access_option_stat_visual(self): spec = mujoco.MjSpec.from_string("""