diff --git a/src/user/user_model.cc b/src/user/user_model.cc index 7402ef91..ea0c0028 100644 --- a/src/user/user_model.cc +++ b/src/user/user_model.cc @@ -892,6 +892,14 @@ mjsElement* mjCModel::NextObject(mjsElement* object, mjtObj type) { } switch (type) { + case mjOBJ_BODY: + case mjOBJ_SITE: + case mjOBJ_GEOM: + case mjOBJ_JOINT: + case mjOBJ_CAMERA: + case mjOBJ_LIGHT: + case mjOBJ_FRAME: + return bodies_[0]->NextChild(object, type, /*recursive=*/true); case mjOBJ_ACTUATOR: return GetNext(actuators_, object); case mjOBJ_SENSOR: diff --git a/src/user/user_objects.cc b/src/user/user_objects.cc index 742f7e3c..370140ab 100644 --- a/src/user/user_objects.cc +++ b/src/user/user_objects.cc @@ -1253,26 +1253,38 @@ mjCBase* mjCBody::FindObject(mjtObj type, std::string _name, bool recursive) { template -static mjsElement* GetNext(std::vector& list, mjsElement* child) { +mjsElement* mjCBody::GetNext(std::vector& list, const mjsElement* child, bool recursive) { + if (list.empty()) { + // no children + return nullptr; + } + if (!child) { - if (list.empty()) { - return nullptr; - } + // first child return list[0]->spec.element; } for (unsigned int i = 0; i < list.size()-1; i++) { + // next child is in this body if (list[i]->spec.element == child) { return list[i+1]->spec.element; } } + + if (recursive && list.back()->spec.element == child) { + // next child is in next body + for (int i=0; i<(int)bodies.size(); i++) { + return bodies[i]->NextChild(NULL, child->elemtype, true); + } + } + return nullptr; } // get next child of given type -mjsElement* mjCBody::NextChild(mjsElement* child, mjtObj type) { +mjsElement* mjCBody::NextChild(mjsElement* child, mjtObj type, bool recursive) { if (type == mjOBJ_UNKNOWN) { if (!child) { throw mjCError(this, "child type must be specified if no child element is given"); @@ -1283,25 +1295,44 @@ mjsElement* mjCBody::NextChild(mjsElement* child, mjtObj type) { throw mjCError(this, "child element is not of requested type"); } + mjsElement* candidate = nullptr; switch (type) { case mjOBJ_BODY: case mjOBJ_XBODY: - return GetNext(bodies, child); + candidate = GetNext(bodies, child, recursive); + break; case mjOBJ_JOINT: - return GetNext(joints, child); + candidate = GetNext(joints, child, recursive); + break; case mjOBJ_GEOM: - return GetNext(geoms, child); + candidate = GetNext(geoms, child, recursive); + break; case mjOBJ_SITE: - return GetNext(sites, child); + candidate = GetNext(sites, child, recursive); + break; case mjOBJ_CAMERA: - return GetNext(cameras, child); + candidate = GetNext(cameras, child, recursive); + break; case mjOBJ_LIGHT: - return GetNext(lights, child); + candidate = GetNext(lights, child, recursive); + break; case mjOBJ_FRAME: - return GetNext(frames, child); + candidate = GetNext(frames, child, recursive); + break; default: - return nullptr; + break; } + + if (!candidate && recursive) { + for (int i=0; i<(int)bodies.size(); i++) { + candidate = bodies[i]->NextChild(child, type, true); + if (candidate) { + return candidate; + } + } + } + + return candidate; } diff --git a/src/user/user_objects.h b/src/user/user_objects.h index 47442f1d..bf514b9b 100644 --- a/src/user/user_objects.h +++ b/src/user/user_objects.h @@ -322,7 +322,7 @@ class mjCBody : public mjCBody_, private mjsBody { const std::vector& get_userdata() { return userdata_; } // get next child of given type - mjsElement* NextChild(mjsElement* child, mjtObj type = mjOBJ_UNKNOWN); + mjsElement* NextChild(mjsElement* child, mjtObj type = mjOBJ_UNKNOWN, bool recursive = false); // reset keyframe references for allowing self-attach void ForgetKeyframes() const; @@ -355,6 +355,10 @@ class mjCBody : public mjCBody_, private mjsBody { template void CopyList(std::vector& dst, const std::vector& src, std::map& fmap, const mjCFrame* pframe = nullptr); + + // gets next child of the same type in this body + template + mjsElement* GetNext(std::vector& list, const mjsElement* child, bool recursive = false); }; diff --git a/test/user/user_api_test.cc b/test/user/user_api_test.cc index 3e1b96e4..bb084d6c 100644 --- a/test/user/user_api_test.cc +++ b/test/user/user_api_test.cc @@ -70,6 +70,7 @@ TEST_F(MujocoTest, TreeTraversal) { mjSpec* spec = mj_makeSpec(); mjsBody* world = mjs_findBody(spec, "world"); mjsBody* body = mjs_addBody(world, 0); + mjsBody* body1 = mjs_addBody(body, 0); mjsSite* site1 = mjs_addSite(body, 0); mjsGeom* geom1 = mjs_addGeom(body, 0); @@ -77,6 +78,7 @@ TEST_F(MujocoTest, TreeTraversal) { mjsSite* site2 = mjs_addSite(body, 0); mjsSite* site3 = mjs_addSite(body, 0); mjsGeom* geom3 = mjs_addGeom(body, 0); + mjsSite* site4 = mjs_addSite(body1, 0); mjs_setString(site1->name, "site1"); mjs_setString(geom1->name, "geom1"); @@ -84,30 +86,40 @@ TEST_F(MujocoTest, TreeTraversal) { mjs_setString(site2->name, "site2"); mjs_setString(site3->name, "site3"); mjs_setString(geom3->name, "geom3"); + mjs_setString(site4->name, "site4"); - mjsElement* a_el1 = mjs_firstElement(spec, mjOBJ_ACTUATOR); - mjsElement* c_el1 = mjs_firstChild(body, mjOBJ_CAMERA); - mjsElement* t_el1 = mjs_firstChild(body, mjOBJ_TENDON); + mjsElement* a_el0 = mjs_firstElement(spec, mjOBJ_ACTUATOR); + mjsElement* l_el0 = mjs_firstElement(spec, mjOBJ_LIGHT); + mjsElement* c_el0 = mjs_firstChild(body, mjOBJ_CAMERA); + mjsElement* t_el0 = mjs_firstChild(body, mjOBJ_TENDON); mjsElement* s_el1 = mjs_firstChild(body, mjOBJ_SITE); mjsElement* s_el2 = mjs_nextChild(body, s_el1); mjsElement* s_el3 = mjs_nextChild(body, s_el2); - mjsElement* s_el4 = mjs_nextChild(body, s_el3); + mjsElement* s_el0 = mjs_nextChild(body, s_el3); + mjsElement* s_el4 = mjs_firstChild(body1, mjOBJ_SITE); mjsElement* g_el1 = mjs_firstChild(body, mjOBJ_GEOM); mjsElement* g_el2 = mjs_nextChild(body, g_el1); mjsElement* g_el3 = mjs_nextChild(body, g_el2); - mjsElement* g_el4 = mjs_nextChild(body, g_el3); + mjsElement* g_el0 = mjs_nextChild(body, g_el3); - EXPECT_EQ(a_el1, nullptr); - EXPECT_EQ(c_el1, nullptr); - EXPECT_EQ(t_el1, nullptr); + EXPECT_EQ(a_el0, nullptr); + EXPECT_EQ(l_el0, nullptr); + EXPECT_EQ(c_el0, nullptr); + EXPECT_EQ(t_el0, nullptr); + EXPECT_EQ(g_el0, nullptr); + EXPECT_EQ(s_el0, nullptr); EXPECT_EQ(s_el1, site1->element); EXPECT_EQ(s_el2, site2->element); EXPECT_EQ(s_el3, site3->element); + EXPECT_EQ(s_el4, site4->element); EXPECT_EQ(g_el1, geom1->element); EXPECT_EQ(g_el2, geom2->element); EXPECT_EQ(g_el3, geom3->element); - EXPECT_EQ(g_el4, nullptr); - EXPECT_EQ(s_el4, nullptr); + EXPECT_EQ(s_el1, mjs_firstElement(spec, mjOBJ_SITE)); + EXPECT_EQ(s_el2, mjs_nextElement(spec, s_el1)); + EXPECT_EQ(s_el3, mjs_nextElement(spec, s_el2)); + EXPECT_EQ(s_el4, mjs_nextElement(spec, s_el3)); + EXPECT_EQ(nullptr, mjs_nextElement(spec, s_el4)); EXPECT_EQ(mjs_findElement(spec, mjOBJ_SITE, "site1"), site1->element); EXPECT_EQ(mjs_findElement(spec, mjOBJ_SITE, "site2"), site2->element); EXPECT_EQ(mjs_findElement(spec, mjOBJ_SITE, "site3"), site3->element);