diff --git a/python/mujoco/specs.cc b/python/mujoco/specs.cc index 2359de42..96bb7baf 100644 --- a/python/mujoco/specs.cc +++ b/python/mujoco/specs.cc @@ -155,6 +155,48 @@ void DefineArray(py::module& m, const std::string& typestr) { }, py::keep_alive<0, 1>(), py::return_value_policy::reference_internal); }; +py::list FindAllImpl(raw::MjsBody& body, mjtObj objtype) { + py::list list; + raw::MjsElement* el = mjs_firstChild(&body, objtype, true); + std::string error = mjs_getError(mjs_getSpec(body.element)); + if (!el && !error.empty()) { + throw pybind11::value_error(error); + } + while (el) { + switch (objtype) { + case mjOBJ_BODY: + list.append(mjs_asBody(el)); + break; + case mjOBJ_CAMERA: + list.append(mjs_asCamera(el)); + break; + case mjOBJ_FRAME: + list.append(mjs_asFrame(el)); + break; + case mjOBJ_GEOM: + list.append(mjs_asGeom(el)); + break; + case mjOBJ_JOINT: + list.append(mjs_asJoint(el)); + break; + case mjOBJ_LIGHT: + list.append(mjs_asLight(el)); + break; + case mjOBJ_SITE: + list.append(mjs_asSite(el)); + break; + default: + // this should never happen + throw pybind11::value_error( + "body.find_all supports the types: body, frame, geom, site, " + "light, camera."); + break; + } + el = mjs_nextChild(&body, el, true); + } + return list; // list of pointers, so they can be copied +} + PYBIND11_MODULE(_specs, m) { auto structs_m = py::module::import("mujoco._structs"); py::function mjmodel_from_spec_ptr = @@ -419,45 +461,31 @@ PYBIND11_MODULE(_specs, m) { mjsBody.def( "find_all", [](raw::MjsBody& self, mjtObj objtype) -> py::list { - py::list list; - raw::MjsElement* el = mjs_firstChild(&self, objtype, true); - std::string error = mjs_getError(mjs_getSpec(self.element)); - if (!el && !error.empty()) { - throw pybind11::value_error(error); + return FindAllImpl(self, objtype); + }, + py::return_value_policy::reference_internal); + mjsBody.def( + "find_all", + [](raw::MjsBody& self, std::string& name) -> py::list { + mjtObj objtype = mjOBJ_UNKNOWN; + if (name == "body") { + objtype = mjOBJ_BODY; + } else if (name == "frame") { + objtype = mjOBJ_FRAME; + } else if (name == "geom") { + objtype = mjOBJ_GEOM; + } else if (name == "site") { + objtype = mjOBJ_SITE; + } else if (name == "light") { + objtype = mjOBJ_LIGHT; + } else if (name == "camera") { + objtype = mjOBJ_CAMERA; + } else { + throw pybind11::value_error( + "body.find_all supports the types: body, frame, geom, site, " + "light, camera."); } - while (el) { - switch (objtype) { - case mjOBJ_BODY: - list.append(mjs_asBody(el)); - break; - case mjOBJ_CAMERA: - list.append(mjs_asCamera(el)); - break; - case mjOBJ_FRAME: - list.append(mjs_asFrame(el)); - break; - case mjOBJ_GEOM: - list.append(mjs_asGeom(el)); - break; - case mjOBJ_JOINT: - list.append(mjs_asJoint(el)); - break; - case mjOBJ_LIGHT: - list.append(mjs_asLight(el)); - break; - case mjOBJ_SITE: - list.append(mjs_asSite(el)); - break; - default: - // this should never happen - throw pybind11::value_error( - "body.find_all supports the types: body, frame, geom, site, " - "light, camera."); - break; - } - el = mjs_nextChild(&self, el, true); - } - return list; + return FindAllImpl(self, objtype); }, py::return_value_policy::reference_internal); mjsBody.def( diff --git a/python/mujoco/specs_test.py b/python/mujoco/specs_test.py index b5955e5c..450c1f78 100644 --- a/python/mujoco/specs_test.py +++ b/python/mujoco/specs_test.py @@ -611,7 +611,6 @@ class SpecsTest(absltest.TestCase): """ spec = mujoco.MjSpec.from_string(main_xml) bodytype = mujoco.mjtObj.mjOBJ_BODY - sitetype = mujoco.mjtObj.mjOBJ_SITE self.assertLen(spec.bodies, 5) self.assertEqual(spec.bodies[1].name, 'body1') self.assertEqual(spec.bodies[2].name, 'body2') @@ -620,23 +619,26 @@ class SpecsTest(absltest.TestCase): self.assertLen(spec.worldbody.find_all(bodytype), 4) self.assertLen(spec.bodies[1].find_all(bodytype), 2) self.assertLen(spec.bodies[3].find_all(bodytype), 1) - self.assertEqual(spec.worldbody.find_all(bodytype)[0].name, 'body1') - self.assertEqual(spec.worldbody.find_all(bodytype)[1].name, 'body2') - self.assertEqual(spec.worldbody.find_all(bodytype)[2].name, 'body3') - self.assertEqual(spec.worldbody.find_all(bodytype)[3].name, 'body4') - self.assertEqual(spec.bodies[1].find_all(bodytype)[0].name, 'body3') - self.assertEqual(spec.bodies[1].find_all(bodytype)[1].name, 'body4') - self.assertEqual(spec.bodies[3].find_all(bodytype)[0].name, 'body4') - self.assertEmpty(spec.bodies[2].find_all(bodytype)) - self.assertEmpty(spec.bodies[4].find_all(bodytype)) - self.assertEqual(spec.worldbody.find_all(sitetype)[0].name, 'site') + self.assertEqual(spec.worldbody.find_all('body')[0].name, 'body1') + self.assertEqual(spec.worldbody.find_all('body')[1].name, 'body2') + self.assertEqual(spec.worldbody.find_all('body')[2].name, 'body3') + self.assertEqual(spec.worldbody.find_all('body')[3].name, 'body4') + self.assertEqual(spec.bodies[1].find_all('body')[0].name, 'body3') + self.assertEqual(spec.bodies[1].find_all('body')[1].name, 'body4') + self.assertEqual(spec.bodies[3].find_all('body')[0].name, 'body4') + self.assertEmpty(spec.bodies[2].find_all('body')) + self.assertEmpty(spec.bodies[4].find_all('body')) + self.assertEqual(spec.worldbody.find_all('site')[0].name, 'site') with self.assertRaises(ValueError) as cm: - spec.worldbody.find_all(mujoco.mjtObj.mjOBJ_ACTUATOR) + spec.worldbody.find_all('actuator') self.assertEqual( str(cm.exception), - 'Error: Body.NextChild supports the types: body, frame, geom, site,' - ' light, camera\nElement name \'world\', id 0', + 'body.find_all supports the types: body, frame, geom, site,' + ' light, camera.', ) + body4 = spec.worldbody.find_all('body')[3] + body4.name = 'body4_new' + self.assertEqual(spec.bodies[4].name, 'body4_new') def test_iterators(self): spec = mujoco.MjSpec()