Add find_all with string input.

PiperOrigin-RevId: 689487165
Change-Id: I48511aec1610b2112f9456e8f410d077c57c7a95
This commit is contained in:
Alessio Quaglino
2024-10-24 13:00:38 -07:00
committed by Copybara-Service
parent a3ea01e57e
commit bebec52869
2 changed files with 82 additions and 52 deletions
+16 -14
View File
@@ -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()