diff --git a/doc/APIreference/functions.rst b/doc/APIreference/functions.rst index 234cbb8b..31bd0d6d 100644 --- a/doc/APIreference/functions.rst +++ b/doc/APIreference/functions.rst @@ -4123,7 +4123,16 @@ mjs_findBody .. mujoco-include:: mjs_findBody -Find body in model by name. +Find body in spec by name. + +.. _mjs_findElement: + +mjs_findElement +~~~~~~~~~~~~~~~ + +.. mujoco-include:: mjs_findElement + +Find element in spec by name. .. _mjs_findChild: @@ -4134,15 +4143,6 @@ mjs_findChild Find child body by name. -.. _mjs_findMesh: - -mjs_findMesh -~~~~~~~~~~~~ - -.. mujoco-include:: mjs_findMesh - -Find mesh by name. - .. _mjs_findFrame: mjs_findFrame @@ -4152,15 +4152,6 @@ mjs_findFrame Find frame by name. -.. _mjs_findKeyframe: - -mjs_findKeyframe -~~~~~~~~~~~~~~~~ - -.. mujoco-include:: mjs_findKeyframe - -Find keyframe by name. - .. _mjs_getDefault: mjs_getDefault diff --git a/doc/changelog.rst b/doc/changelog.rst index 3ca2baab..ab3a708e 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -15,6 +15,8 @@ General :ref:`ccd_tolerance` and :ref:`ccd_iterations`, both in XML and in the :ref:`mjOption` struct. This is because the new convex collision detection pipeline (see below) does not use the MPR algorithm. The semantics of these options remain identical. + - The functions ``mjs_findMesh`` and ``mjs_findKeyframe`` were replaced by ``mjs_findElement``, which allows to look + for any object type. - Added a new way of defining :ref:`connect` equality constraints, using two sites rather than bodies. The new semantic is useful when the assumption that the constraint is satisfied in the base configuration does not diff --git a/doc/includes/references.h b/doc/includes/references.h index b6de8a42..97ead8a2 100644 --- a/doc/includes/references.h +++ b/doc/includes/references.h @@ -3572,10 +3572,9 @@ mjsTexture* mjs_addTexture(mjSpec* s); mjsMaterial* mjs_addMaterial(mjSpec* s, mjsDefault* def); mjSpec* mjs_getSpec(mjsBody* body); mjsBody* mjs_findBody(mjSpec* s, const char* name); +mjsElement* mjs_findElement(mjSpec* s, mjtObj type, const char* name); mjsBody* mjs_findChild(mjsBody* body, const char* name); -mjsMesh* mjs_findMesh(mjSpec* s, const char* name); mjsFrame* mjs_findFrame(mjSpec* s, const char* name); -mjsKey* mjs_findKeyframe(mjSpec* s, const char* name); mjsDefault* mjs_getDefault(mjsElement* element); mjsDefault* mjs_findDefault(mjSpec* s, const char* classname); mjsDefault* mjs_getSpecDefault(mjSpec* s); diff --git a/include/mujoco/mujoco.h b/include/mujoco/mujoco.h index 98c415e1..8c6ffa17 100644 --- a/include/mujoco/mujoco.h +++ b/include/mujoco/mujoco.h @@ -1522,21 +1522,18 @@ MJAPI mjsMaterial* mjs_addMaterial(mjSpec* s, mjsDefault* def); // Get spec from body. MJAPI mjSpec* mjs_getSpec(mjsBody* body); -// Find body in model by name. +// Find body in spec by name. MJAPI mjsBody* mjs_findBody(mjSpec* s, const char* name); +// Find element in spec by name. +MJAPI mjsElement* mjs_findElement(mjSpec* s, mjtObj type, const char* name); + // Find child body by name. MJAPI mjsBody* mjs_findChild(mjsBody* body, const char* name); -// Find mesh by name. -MJAPI mjsMesh* mjs_findMesh(mjSpec* s, const char* name); - // Find frame by name. MJAPI mjsFrame* mjs_findFrame(mjSpec* s, const char* name); -// Find keyframe by name. -MJAPI mjsKey* mjs_findKeyframe(mjSpec* s, const char* name); - // Get default corresponding to an element. MJAPI mjsDefault* mjs_getDefault(mjsElement* element); diff --git a/introspect/functions.py b/introspect/functions.py index c90ddb7e..e6358821 100644 --- a/introspect/functions.py +++ b/introspect/functions.py @@ -9678,7 +9678,33 @@ FUNCTIONS: Mapping[str, FunctionDecl] = dict([ ), ), ), - doc='Find body in model by name.', + doc='Find body in spec by name.', + )), + ('mjs_findElement', + FunctionDecl( + name='mjs_findElement', + return_type=PointerType( + inner_type=ValueType(name='mjsElement'), + ), + parameters=( + FunctionParameterDecl( + name='s', + type=PointerType( + inner_type=ValueType(name='mjSpec'), + ), + ), + FunctionParameterDecl( + name='type', + type=ValueType(name='mjtObj'), + ), + FunctionParameterDecl( + name='name', + type=PointerType( + inner_type=ValueType(name='char', is_const=True), + ), + ), + ), + doc='Find element in spec by name.', )), ('mjs_findChild', FunctionDecl( @@ -9702,28 +9728,6 @@ FUNCTIONS: Mapping[str, FunctionDecl] = dict([ ), doc='Find child body by name.', )), - ('mjs_findMesh', - FunctionDecl( - name='mjs_findMesh', - return_type=PointerType( - inner_type=ValueType(name='mjsMesh'), - ), - parameters=( - FunctionParameterDecl( - name='s', - type=PointerType( - inner_type=ValueType(name='mjSpec'), - ), - ), - FunctionParameterDecl( - name='name', - type=PointerType( - inner_type=ValueType(name='char', is_const=True), - ), - ), - ), - doc='Find mesh by name.', - )), ('mjs_findFrame', FunctionDecl( name='mjs_findFrame', @@ -9746,28 +9750,6 @@ FUNCTIONS: Mapping[str, FunctionDecl] = dict([ ), doc='Find frame by name.', )), - ('mjs_findKeyframe', - FunctionDecl( - name='mjs_findKeyframe', - return_type=PointerType( - inner_type=ValueType(name='mjsKey'), - ), - parameters=( - FunctionParameterDecl( - name='s', - type=PointerType( - inner_type=ValueType(name='mjSpec'), - ), - ), - FunctionParameterDecl( - name='name', - type=PointerType( - inner_type=ValueType(name='char', is_const=True), - ), - ), - ), - doc='Find keyframe by name.', - )), ('mjs_getDefault', FunctionDecl( name='mjs_getDefault', diff --git a/python/mujoco/specs.cc b/python/mujoco/specs.cc index 0f4ab874..27e7d819 100644 --- a/python/mujoco/specs.cc +++ b/python/mujoco/specs.cc @@ -133,24 +133,12 @@ PYBIND11_MODULE(_specs, m) { return mjs_findBody(self.ptr, name.c_str()); }, py::return_value_policy::reference_internal); - mjSpec.def( - "find_mesh", - [](MjSpec& self, std::string& name) -> raw::MjsMesh* { - return mjs_findMesh(self.ptr, name.c_str()); - }, - py::return_value_policy::reference_internal); mjSpec.def( "find_frame", [](MjSpec& self, std::string& name) -> raw::MjsFrame* { return mjs_findFrame(self.ptr, name.c_str()); }, py::return_value_policy::reference_internal); - mjSpec.def( - "find_keyframe", - [](MjSpec& self, std::string& name) -> raw::MjsKey* { - return mjs_findKeyframe(self.ptr, name.c_str()); - }, - py::return_value_policy::reference_internal); mjSpec.def( "find_default", [](MjSpec& self, std::string& classname) -> raw::MjsDefault* { diff --git a/src/user/user_api.cc b/src/user/user_api.cc index 228299f1..4253f2ed 100644 --- a/src/user/user_api.cc +++ b/src/user/user_api.cc @@ -542,18 +542,34 @@ mjsDefault* mjs_getSpecDefault(mjSpec* s) { // find body in model by name mjsBody* mjs_findBody(mjSpec* s, const char* name) { - mjCModel* model = static_cast(s->element); - mjCBase* body = 0; - if (model->IsCompiled()) { - body = model->FindObject(mjOBJ_BODY, std::string(name)); // fast lookup - } else { - body = model->FindBody(model->GetWorld(), std::string(name)); // recursive search - } + mjsElement* body = mjs_findElement(s, mjOBJ_BODY, name); return body ? &(static_cast(body)->spec) : nullptr; } +// find element in spec by name +mjsElement* mjs_findElement(mjSpec* s, mjtObj type, const char* name) { + mjCModel* model = static_cast(s->element); + if (model->IsCompiled()) { + return model->FindObject(type, std::string(name)); // fast lookup + } + switch (type) { + case mjOBJ_BODY: + case mjOBJ_SITE: + case mjOBJ_GEOM: + case mjOBJ_JOINT: + case mjOBJ_CAMERA: + case mjOBJ_LIGHT: + case mjOBJ_FRAME: + return model->FindTree(model->GetWorld(), type, std::string(name)); // recursive search + default: + return model->FindObject(type, std::string(name)); // always available + } +} + + + // find child of a body by name mjsBody* mjs_findChild(mjsBody* bodyspec, const char* name) { mjCBody* body = static_cast(bodyspec->element); @@ -563,33 +579,14 @@ mjsBody* mjs_findChild(mjsBody* bodyspec, const char* name) { -// find mesh by name -mjsMesh* mjs_findMesh(mjSpec* s, const char* name) { - mjCModel* model = static_cast(s->element); - mjCMesh* mesh = (mjCMesh*)model->FindObject(mjOBJ_MESH, std::string(name)); - return mesh ? &(static_cast(mesh)->spec) : nullptr; -} - - - // find frame by name mjsFrame* mjs_findFrame(mjSpec* s, const char* name) { - mjCModel* model = static_cast(s->element); - mjCFrame* frame = (mjCFrame*)model->FindFrame(model->GetWorld(), std::string(name)); + mjsElement* frame = mjs_findElement(s, mjOBJ_FRAME, name); return frame ? &(static_cast(frame)->spec) : nullptr; } -// find keyframe by name -mjsKey* mjs_findKeyframe(mjSpec* s, const char* name) { - mjCModel* model = static_cast(s->element); - mjCKey* key = (mjCKey*)model->FindObject(mjOBJ_KEY, std::string(name)); - return key ? &(static_cast(key)->spec) : nullptr; -} - - - // set frame void mjs_setFrame(mjsElement* dest, mjsFrame* frame) { if (!frame) { diff --git a/src/user/user_api.h b/src/user/user_api.h index 82ce193a..35799bad 100644 --- a/src/user/user_api.h +++ b/src/user/user_api.h @@ -188,21 +188,18 @@ MJAPI mjSpec* mjs_getSpec(mjsBody* body); // Find spec (model asset) by name. MJAPI mjSpec* mjs_findSpec(mjSpec* spec, const char* name); -// Find body in model by name. +// Find body in spec by name. MJAPI mjsBody* mjs_findBody(mjSpec* s, const char* name); +// Find element in spec by name. +MJAPI mjsElement* mjs_findElement(mjSpec* s, mjtObj type, const char* name); + // Find child body by name. MJAPI mjsBody* mjs_findChild(mjsBody* body, const char* name); -// Find mesh by name. -MJAPI mjsMesh* mjs_findMesh(mjSpec* s, const char* name); - // Find frame by name. MJAPI mjsFrame* mjs_findFrame(mjSpec* s, const char* name); -// Find keyframe by name. -MJAPI mjsKey* mjs_findKeyframe(mjSpec* s, const char* name); - // Get default corresponding to an element. MJAPI mjsDefault* mjs_getDefault(mjsElement* element); diff --git a/src/user/user_model.cc b/src/user/user_model.cc index 582c932c..7402ef91 100644 --- a/src/user/user_model.cc +++ b/src/user/user_model.cc @@ -1031,33 +1031,61 @@ mjCBase* mjCModel::FindObject(mjtObj type, string name) const { // find body by name -mjCBody* mjCModel::FindBody(mjCBody* body, std::string name) { - if (body->name == name) { - return body; +mjCBase* mjCModel::FindTree(mjCBody* body, mjtObj type, std::string name) { + switch (type) { + case mjOBJ_BODY: + if (body->name == name) { + return body; + } + break; + case mjOBJ_SITE: + for (auto site : body->sites) { + if (site->name == name) { + return site; + } + } + break; + case mjOBJ_GEOM: + for (auto geom : body->geoms) { + if (geom->name == name) { + return geom; + } + } + break; + case mjOBJ_JOINT: + for (auto joint : body->joints) { + if (joint->name == name) { + return joint; + } + } + break; + case mjOBJ_CAMERA: + for (auto camera : body->cameras) { + if (camera->name == name) { + return camera; + } + } + break; + case mjOBJ_LIGHT: + for (auto light : body->lights) { + if (light->name == name) { + return light; + } + } + break; + case mjOBJ_FRAME: + for (auto frame : body->frames) { + if (frame->name == name) { + return frame; + } + } + break; + default: + return nullptr; } for (auto child : body->bodies) { - auto candidate = FindBody(child, name); - if (candidate) { - return candidate; - } - } - - return nullptr; -} - - - -// find frame by name -mjCFrame* mjCModel::FindFrame(mjCBody* body, std::string name) const{ - for (auto frame : body->frames) { - if (frame->name == name) { - return frame; - } - } - - for (auto body : body->bodies) { - auto candidate = FindFrame(body, name); + auto candidate = FindTree(child, type, name); if (candidate) { return candidate; } diff --git a/src/user/user_model.h b/src/user/user_model.h index 026d96ed..9a0e7939 100644 --- a/src/user/user_model.h +++ b/src/user/user_model.h @@ -223,16 +223,15 @@ class mjCModel : public mjCModel_, private mjSpec { mjsElement* NextObject(mjsElement* object, mjtObj type = mjOBJ_UNKNOWN); // next object of specified type // API for access to other variables - bool IsCompiled() const; // is model already compiled - const mjCError& GetError() const; // get reference of error object - void SetError(const mjCError& error) { errInfo = error; } // set value of error object - mjCBody* GetWorld(); // pointer to world body - mjCDef* FindDefault(std::string name); // find defaults class name - mjCDef* AddDefault(std::string name, mjCDef* parent = nullptr); // add defaults class to array - mjCBase* FindObject(mjtObj type, std::string name) const; // find object given type and name - mjCBody* FindBody(mjCBody* body, std::string name); // find body given name - mjCFrame* FindFrame(mjCBody* body, std::string name) const; // find frame given name - mjSpec* FindSpec(std::string name) const; // find spec given name + bool IsCompiled() const; // is model already compiled + const mjCError& GetError() const; // get reference of error object + void SetError(const mjCError& error) { errInfo = error; } // set value of error object + mjCBody* GetWorld(); // pointer to world body + mjCDef* FindDefault(std::string name); // find defaults class name + mjCDef* AddDefault(std::string name, mjCDef* parent = nullptr); // add defaults class to array + mjCBase* FindObject(mjtObj type, std::string name) const; // find object given type and name + mjCBase* FindTree(mjCBody* body, mjtObj type, std::string name); // find tree object given name + mjSpec* FindSpec(std::string name) const; // find spec given name void SetActivePlugins(const std::vector>&& active_plugins) { active_plugins_ = std::move(active_plugins); } diff --git a/src/user/user_objects.cc b/src/user/user_objects.cc index 956f7a03..742f7e3c 100644 --- a/src/user/user_objects.cc +++ b/src/user/user_objects.cc @@ -1408,8 +1408,8 @@ void mjCBody::ForgetKeyframes() const { joint->qpos_.clear(); joint->qvel_.clear(); } - model->FindBody((mjCBody*)this, name)->mpos_.clear(); // this is a hack to avoid const - model->FindBody((mjCBody*)this, name)->mquat_.clear(); // this is a hack to avoid const + ((mjCBody*)this)->mpos_.clear(); + ((mjCBody*)this)->mquat_.clear(); for (auto body : bodies) { body->ForgetKeyframes(); } diff --git a/src/xml/xml_urdf.cc b/src/xml/xml_urdf.cc index b623f79c..2c2875c9 100644 --- a/src/xml/xml_urdf.cc +++ b/src/xml/xml_urdf.cc @@ -592,7 +592,7 @@ mjsGeom* mjXURDF::Geom(XMLElement* geom_elem, mjsBody* pbody, bool collision) { meshname = mjuu_stripext(meshname); // look for existing mesh - mjsMesh* mesh = mjs_findMesh(spec, meshname.c_str()); + mjsMesh* mesh = mjs_asMesh(mjs_findElement(spec, mjOBJ_MESH, meshname.c_str())); mjsMesh* pmesh = 0; // does not exist: create diff --git a/test/user/user_api_test.cc b/test/user/user_api_test.cc index cc88aec9..3e1b96e4 100644 --- a/test/user/user_api_test.cc +++ b/test/user/user_api_test.cc @@ -78,6 +78,13 @@ TEST_F(MujocoTest, TreeTraversal) { mjsSite* site3 = mjs_addSite(body, 0); mjsGeom* geom3 = mjs_addGeom(body, 0); + mjs_setString(site1->name, "site1"); + mjs_setString(geom1->name, "geom1"); + mjs_setString(geom2->name, "geom2"); + mjs_setString(site2->name, "site2"); + mjs_setString(site3->name, "site3"); + mjs_setString(geom3->name, "geom3"); + mjsElement* a_el1 = mjs_firstElement(spec, mjOBJ_ACTUATOR); mjsElement* c_el1 = mjs_firstChild(body, mjOBJ_CAMERA); mjsElement* t_el1 = mjs_firstChild(body, mjOBJ_TENDON); @@ -101,6 +108,12 @@ TEST_F(MujocoTest, TreeTraversal) { EXPECT_EQ(g_el3, geom3->element); EXPECT_EQ(g_el4, nullptr); EXPECT_EQ(s_el4, nullptr); + 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); + EXPECT_EQ(mjs_findElement(spec, mjOBJ_GEOM, "geom1"), geom1->element); + EXPECT_EQ(mjs_findElement(spec, mjOBJ_GEOM, "geom2"), geom2->element); + EXPECT_EQ(mjs_findElement(spec, mjOBJ_GEOM, "geom3"), geom3->element); mj_deleteSpec(spec); }