diff --git a/python/mujoco/specs_test.py b/python/mujoco/specs_test.py index 257bdd6b..e950e97a 100644 --- a/python/mujoco/specs_test.py +++ b/python/mujoco/specs_test.py @@ -642,68 +642,56 @@ class SpecsTest(absltest.TestCase): + - - - - - - - - + + + + + + + + - """ spec = mujoco.MjSpec.from_string(main_xml) bodytype = mujoco.mjtObj.mjOBJ_BODY - self.assertLen(spec.bodies, 5) + self.assertLen(spec.bodies, 4) self.assertLen(spec.sites, 5) - self.assertLen(spec.worldbody.find_all('body'), 4) + self.assertLen(spec.worldbody.find_all('body'), 3) self.assertLen(spec.worldbody.find_all('site'), 5) - self.assertLen(spec.worldbody.find_all('joint'), 1) - self.assertLen(spec.worldbody.find_all('geom'), 1) + self.assertLen(spec.worldbody.find_all('geom'), 3) self.assertEqual(spec.bodies[1].name, 'body1') self.assertEqual(spec.bodies[2].name, 'body2') self.assertEqual(spec.bodies[3].name, 'body3') - self.assertEqual(spec.bodies[4].name, 'body4') - self.assertEqual(spec.bodies[1].parent, spec.worldbody) - self.assertEqual(spec.bodies[2].parent, spec.worldbody) + self.assertEqual(spec.bodies[1].parent, spec.bodies[0]) + self.assertEqual(spec.bodies[2].parent, spec.bodies[0]) self.assertEqual(spec.bodies[3].parent, spec.bodies[1]) - self.assertEqual(spec.bodies[4].parent, spec.bodies[3]) - 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('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.assertLen(spec.worldbody.find_all(bodytype), 3) + self.assertLen(spec.bodies[1].find_all(bodytype), 1) + self.assertEmpty(spec.bodies[3].find_all(bodytype)) 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[3].find_all('body')) 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, 'site1') - self.assertEqual(spec.worldbody.find_all('site')[1].name, 'site2') - self.assertEqual(spec.worldbody.find_all('site')[2].name, 'site3') - self.assertEqual(spec.worldbody.find_all('site')[3].name, 'site4') - self.assertEqual(spec.worldbody.find_all('site')[4].name, 'site5') - self.assertEmpty(spec.bodies[2].sites) - self.assertLen(spec.bodies[3].sites, 4) - self.assertLen(spec.bodies[4].sites, 1) - self.assertEqual(spec.bodies[3].sites[0].name, 'site1') - self.assertEqual(spec.bodies[3].sites[1].name, 'site2') - self.assertEqual(spec.bodies[3].sites[2].name, 'site3') - self.assertEqual(spec.bodies[3].sites[3].name, 'site4') - self.assertEqual(spec.bodies[4].sites[0].name, 'site5') - self.assertEqual(spec.bodies[3].sites[0].parent, spec.bodies[3]) - self.assertEqual(spec.bodies[3].sites[1].parent, spec.bodies[3]) - self.assertEqual(spec.bodies[3].sites[2].parent, spec.bodies[3]) - self.assertEqual(spec.bodies[3].sites[3].parent, spec.bodies[3]) - self.assertEqual(spec.bodies[4].sites[0].parent, spec.bodies[4]) + for i, body in enumerate(spec.worldbody.find_all('body')): + self.assertEqual(body.name, 'body' + str(i + 1)) + for i, site in enumerate(spec.worldbody.find_all('site')): + self.assertEqual(site.name, 'site' + str(i + 1)) + self.assertLen(spec.bodies[1].sites, 3) + self.assertLen(spec.bodies[2].sites, 1) + self.assertLen(spec.bodies[3].sites, 1) + self.assertEqual(spec.bodies[1].sites[0].name, 'site1') + self.assertEqual(spec.bodies[1].sites[1].name, 'site2') + self.assertEqual(spec.bodies[1].sites[2].name, 'site3') + self.assertEqual(spec.bodies[3].sites[0].name, 'site4') + self.assertEqual(spec.bodies[2].sites[0].name, 'site5') + for body in spec.bodies: + for site in body.sites: + self.assertEqual(site.parent, body) with self.assertRaises(ValueError) as cm: spec.worldbody.find_all('actuator') self.assertEqual( @@ -711,9 +699,9 @@ class SpecsTest(absltest.TestCase): 'body.find_all supports the types: body, frame, geom, site,' ' joint, light, camera.', ) - body4 = spec.worldbody.find_all('body')[3] - body4.name = 'body4_new' - self.assertEqual(spec.bodies[4].name, 'body4_new') + body3 = spec.worldbody.find_all('body')[2] + body3.name = 'body3_new' + self.assertEqual(spec.bodies[3].name, 'body3_new') def test_geom_list(self): main_xml = """ diff --git a/test/user/user_api_test.cc b/test/user/user_api_test.cc index 73160585..1b838467 100644 --- a/test/user/user_api_test.cc +++ b/test/user/user_api_test.cc @@ -72,18 +72,18 @@ TEST_F(MujocoTest, TreeTraversal) { static constexpr char xml[] = R"( - - + + + - - + @@ -95,9 +95,9 @@ TEST_F(MujocoTest, TreeTraversal) { ASSERT_THAT(spec, NotNull()) << err.data(); mjsBody* world = mjs_findBody(spec, "world"); - mjsBody* body = mjs_findBody(spec, "body"); mjsBody* body1 = mjs_findBody(spec, "body1"); mjsBody* body2 = mjs_findBody(spec, "body2"); + mjsBody* body3 = mjs_findBody(spec, "body3"); mjsElement* site1 = mjs_findElement(spec, mjOBJ_SITE, "site1"); mjsElement* site2 = mjs_findElement(spec, mjOBJ_SITE, "site2"); mjsElement* site3 = mjs_findElement(spec, mjOBJ_SITE, "site3"); @@ -110,30 +110,30 @@ TEST_F(MujocoTest, TreeTraversal) { // test nonexistent EXPECT_EQ(mjs_firstElement(spec, mjOBJ_ACTUATOR), nullptr); EXPECT_EQ(mjs_firstElement(spec, mjOBJ_LIGHT), nullptr); - EXPECT_EQ(mjs_firstChild(body, mjOBJ_CAMERA, /*recurse=*/true), nullptr); - EXPECT_EQ(mjs_firstChild(body, mjOBJ_TENDON, /*recurse=*/true), nullptr); + EXPECT_EQ(mjs_firstChild(body1, mjOBJ_CAMERA, /*recurse=*/true), nullptr); + EXPECT_EQ(mjs_firstChild(body1, mjOBJ_TENDON, /*recurse=*/true), nullptr); // test first, nonrecursive EXPECT_EQ(mjs_firstElement(spec, mjOBJ_BODY), world->element); EXPECT_EQ(site1, mjs_firstElement(spec, mjOBJ_SITE)); - EXPECT_EQ(site1, mjs_firstChild(body, mjOBJ_SITE, /*recurse=*/false)); - EXPECT_EQ(geom1, mjs_firstChild(body, mjOBJ_GEOM, /*recurse=*/false)); - EXPECT_EQ(site4, mjs_firstChild(body1, mjOBJ_SITE, /*recurse=*/false)); - EXPECT_EQ(site5, mjs_firstChild(body2, mjOBJ_SITE, /*recurse=*/false)); + EXPECT_EQ(site1, mjs_firstChild(body1, mjOBJ_SITE, /*recurse=*/false)); + EXPECT_EQ(geom1, mjs_firstChild(body1, mjOBJ_GEOM, /*recurse=*/false)); + EXPECT_EQ(site4, mjs_firstChild(body2, mjOBJ_SITE, /*recurse=*/false)); + EXPECT_EQ(site5, mjs_firstChild(body3, mjOBJ_SITE, /*recurse=*/false)); // test first, recursive EXPECT_EQ(site1, mjs_firstChild(world, mjOBJ_SITE, /*recurse=*/true)); EXPECT_EQ(geom1, mjs_firstChild(world, mjOBJ_GEOM, /*recurse=*/true)); - EXPECT_EQ(site4, mjs_firstChild(body1, mjOBJ_SITE, /*recurse=*/true)); - EXPECT_EQ(site5, mjs_firstChild(body2, mjOBJ_SITE, /*recurse=*/true)); + EXPECT_EQ(site4, mjs_firstChild(body2, mjOBJ_SITE, /*recurse=*/true)); + EXPECT_EQ(site5, mjs_firstChild(body3, mjOBJ_SITE, /*recurse=*/true)); // text next, nonrecursive - EXPECT_EQ(site2, mjs_nextChild(body, site1, /*recursive=*/false)); - EXPECT_EQ(site3, mjs_nextChild(body, site2, /*recursive=*/false)); - EXPECT_EQ(nullptr, mjs_nextChild(body, site3, /*recursive=*/false)); - EXPECT_EQ(geom2, mjs_nextChild(body, geom1, /*recursive=*/false)); - EXPECT_EQ(geom3, mjs_nextChild(body, geom2, /*recursive=*/false)); - EXPECT_EQ(nullptr, mjs_nextChild(body, geom3, /*recursive=*/false)); + EXPECT_EQ(site2, mjs_nextChild(body1, site1, /*recursive=*/false)); + EXPECT_EQ(site3, mjs_nextChild(body1, site2, /*recursive=*/false)); + EXPECT_EQ(nullptr, mjs_nextChild(body1, site3, /*recursive=*/false)); + EXPECT_EQ(geom2, mjs_nextChild(body1, geom1, /*recursive=*/false)); + EXPECT_EQ(geom3, mjs_nextChild(body1, geom2, /*recursive=*/false)); + EXPECT_EQ(nullptr, mjs_nextChild(body1, geom3, /*recursive=*/false)); EXPECT_EQ(mjs_nextElement(spec, site1), site2); EXPECT_EQ(mjs_nextElement(spec, site2), site3); EXPECT_EQ(mjs_nextElement(spec, site3), site4); @@ -144,15 +144,15 @@ TEST_F(MujocoTest, TreeTraversal) { EXPECT_EQ(mjs_nextElement(spec, geom3), nullptr); // text next, recursive - EXPECT_EQ(site2, mjs_nextChild(body, site1, /*recursive=*/true)); - EXPECT_EQ(site3, mjs_nextChild(body, site2, /*recursive=*/true)); - EXPECT_EQ(site4, mjs_nextChild(body, site3, /*recursive=*/true)); + EXPECT_EQ(site2, mjs_nextChild(body1, site1, /*recursive=*/true)); + EXPECT_EQ(site3, mjs_nextChild(body1, site2, /*recursive=*/true)); + EXPECT_EQ(site4, mjs_nextChild(body1, site3, /*recursive=*/true)); EXPECT_EQ(site4, mjs_nextChild(world, site3, /*recursive=*/true)); EXPECT_EQ(site5, mjs_nextChild(world, site4, /*recursive=*/true)); - EXPECT_EQ(nullptr, mjs_nextChild(body, site5, /*recursive=*/true)); - EXPECT_EQ(geom2, mjs_nextChild(body, geom1, /*recursive=*/true)); - EXPECT_EQ(geom3, mjs_nextChild(body, geom2, /*recursive=*/true)); - EXPECT_EQ(nullptr, mjs_nextChild(body, geom3, /*recursive=*/true)); + EXPECT_EQ(nullptr, mjs_nextChild(body1, site5, /*recursive=*/true)); + EXPECT_EQ(geom2, mjs_nextChild(body1, geom1, /*recursive=*/true)); + EXPECT_EQ(geom3, mjs_nextChild(body1, geom2, /*recursive=*/true)); + EXPECT_EQ(nullptr, mjs_nextChild(body1, geom3, /*recursive=*/true)); // check compilation ordering of sites mjModel* model = mj_compile(spec, nullptr);