From c198a69bbf7eb53ed856a0ae9c2475e1bd603ff2 Mon Sep 17 00:00:00 2001 From: Alessio Quaglino Date: Thu, 18 Jul 2024 14:33:49 -0700 Subject: [PATCH] Return null pointer from mjs_firstChild and mjs_firstElement when the requested list is empty. Fixes #1824. PiperOrigin-RevId: 653758711 Change-Id: I6ede762cafd88dcbb0dfc9c2999b7aee24bd2e2a --- src/user/user_model.cc | 39 ++++++++++++++++++++++---------------- src/user/user_objects.cc | 21 +++++++++++++------- test/user/user_api_test.cc | 4 ++++ 3 files changed, 41 insertions(+), 23 deletions(-) diff --git a/src/user/user_model.cc b/src/user/user_model.cc index a810f6c3..a096bf7c 100644 --- a/src/user/user_model.cc +++ b/src/user/user_model.cc @@ -671,6 +671,13 @@ mjCBase* mjCModel::GetObject(mjtObj type, int id) { template static mjsElement* GetNext(std::vector& list, mjsElement* child) { + if (!child) { + if (list.empty()) { + return nullptr; + } + return list[0]->spec.element; + } + // TODO: use id for direct indexing instead of a loop for (unsigned int i = 0; i < list.size()-1; i++) { if (list[i]->spec.element == child) { @@ -696,37 +703,37 @@ mjsElement* mjCModel::NextObject(mjsElement* object, mjtObj type) { switch (type) { case mjOBJ_ACTUATOR: - return object ? GetNext(actuators_, object) : actuators_[0]; + return GetNext(actuators_, object); case mjOBJ_SENSOR: - return object ? GetNext(sensors_, object) : sensors_[0]; + return GetNext(sensors_, object); case mjOBJ_FLEX: - return object ? GetNext(flexes_, object) : flexes_[0]; + return GetNext(flexes_, object); case mjOBJ_PAIR: - return object ? GetNext(pairs_, object) : pairs_[0]; + return GetNext(pairs_, object); case mjOBJ_EXCLUDE: - return object ? GetNext(excludes_, object) : excludes_[0]; + return GetNext(excludes_, object); case mjOBJ_EQUALITY: - return object ? GetNext(equalities_, object) : equalities_[0]; + return GetNext(equalities_, object); case mjOBJ_TENDON: - return object ? GetNext(tendons_, object) : tendons_[0]; + return GetNext(tendons_, object); case mjOBJ_NUMERIC: - return object ? GetNext(numerics_, object) : numerics_[0]; + return GetNext(numerics_, object); case mjOBJ_TEXT: - return object ? GetNext(texts_, object) : texts_[0]; + return GetNext(texts_, object); case mjOBJ_TUPLE: - return object ? GetNext(tuples_, object) : tuples_[0]; + return GetNext(tuples_, object); case mjOBJ_KEY: - return object ? GetNext(keys_, object) : keys_[0]; + return GetNext(keys_, object); case mjOBJ_MESH: - return object ? GetNext(meshes_, object) : meshes_[0]; + return GetNext(meshes_, object); case mjOBJ_HFIELD: - return object ? GetNext(hfields_, object) : hfields_[0]; + return GetNext(hfields_, object); case mjOBJ_SKIN: - return object ? GetNext(skins_, object) : skins_[0]; + return GetNext(skins_, object); case mjOBJ_TEXTURE: - return object ? GetNext(textures_, object) : textures_[0]; + return GetNext(textures_, object); case mjOBJ_MATERIAL: - return object ? GetNext(materials_, object) : materials_[0]; + return GetNext(materials_, object); default: return nullptr; } diff --git a/src/user/user_objects.cc b/src/user/user_objects.cc index df3980a0..5d9b670f 100644 --- a/src/user/user_objects.cc +++ b/src/user/user_objects.cc @@ -1250,6 +1250,13 @@ mjCBase* mjCBody::FindObject(mjtObj type, std::string _name, bool recursive) { template static mjsElement* GetNext(std::vector& list, mjsElement* child) { + if (!child) { + if (list.empty()) { + return nullptr; + } + return list[0]->spec.element; + } + for (unsigned int i = 0; i < list.size()-1; i++) { if (list[i]->spec.element == child) { return list[i+1]->spec.element; @@ -1275,19 +1282,19 @@ mjsElement* mjCBody::NextChild(mjsElement* child, mjtObj type) { switch (type) { case mjOBJ_BODY: case mjOBJ_XBODY: - return child ? GetNext(bodies, child) : bodies[0]; + return GetNext(bodies, child); case mjOBJ_JOINT: - return child ? GetNext(joints, child) : joints[0]; + return GetNext(joints, child); case mjOBJ_GEOM: - return child ? GetNext(geoms, child) : geoms[0]; + return GetNext(geoms, child); case mjOBJ_SITE: - return child ? GetNext(sites, child) : sites[0]; + return GetNext(sites, child); case mjOBJ_CAMERA: - return child ? GetNext(cameras, child) : cameras[0]; + return GetNext(cameras, child); case mjOBJ_LIGHT: - return child ? GetNext(lights, child) : lights[0]; + return GetNext(lights, child); case mjOBJ_FRAME: - return child ? GetNext(frames, child) : frames[0]; + return GetNext(frames, child); default: return nullptr; } diff --git a/test/user/user_api_test.cc b/test/user/user_api_test.cc index bebdc010..1e811104 100644 --- a/test/user/user_api_test.cc +++ b/test/user/user_api_test.cc @@ -75,6 +75,8 @@ TEST_F(MujocoTest, TreeTraversal) { mjsSite* site3 = mjs_addSite(body, 0); mjsGeom* geom3 = mjs_addGeom(body, 0); + 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* s_el1 = mjs_firstChild(body, mjOBJ_SITE); mjsElement* s_el2 = mjs_nextChild(body, s_el1); @@ -85,6 +87,8 @@ TEST_F(MujocoTest, TreeTraversal) { mjsElement* g_el3 = mjs_nextChild(body, g_el2); mjsElement* g_el4 = mjs_nextChild(body, g_el3); + EXPECT_EQ(a_el1, nullptr); + EXPECT_EQ(c_el1, nullptr); EXPECT_EQ(t_el1, nullptr); EXPECT_EQ(s_el1, site1->element); EXPECT_EQ(s_el2, site2->element);