From ff39d0b8120c4f93525508d8f80b1e12234f199f Mon Sep 17 00:00:00 2001 From: Kyle Bayes Date: Mon, 12 Feb 2024 03:53:56 -0800 Subject: [PATCH] Keep include elements while parsing XML. No change in behavior. PiperOrigin-RevId: 606203141 Change-Id: Ica06f7c121a90597cd29f74692230365edd4302f --- src/xml/xml.cc | 206 +++++++++++----------- src/xml/xml_native_reader.cc | 194 ++++++++++---------- src/xml/xml_util.cc | 274 ++++++++++++++--------------- src/xml/xml_util.h | 32 ++-- test/xml/xml_native_reader_test.cc | 119 ++++++++++++- 5 files changed, 460 insertions(+), 365 deletions(-) diff --git a/src/xml/xml.cc b/src/xml/xml.cc index 36710b44..706a0e7c 100644 --- a/src/xml/xml.cc +++ b/src/xml/xml.cc @@ -20,25 +20,27 @@ #include #endif +#include #include +#include + +#include "tinyxml2.h" #include +#include #include "cc/array_safety.h" #include "engine/engine_crossplatform.h" #include "engine/engine_resource.h" #include "engine/engine_vfs.h" #include "user/user_model.h" -#include "user/user_util.h" #include "xml/xml_native_reader.h" #include "xml/xml_native_writer.h" #include "xml/xml_urdf.h" #include "xml/xml_util.h" -#include "tinyxml2.h" namespace { using std::string; -using std::vector; using tinyxml2::XMLDocument; using tinyxml2::XMLElement; using tinyxml2::XMLNode; @@ -111,106 +113,101 @@ string mjWriteXML(mjCModel* model, char* error, int error_sz) { // find include elements recursively, replace them with subtree from xml file -static XMLElement* mjIncludeXML(XMLElement* elem, string dir, - const mjVFS* vfs, vector& included) { - // include element: process - if (!strcasecmp(elem->Value(), "include")) { - // make sure include has no children - if (!elem->NoChildren()) { - throw mjXError(elem, "Include element cannot have children"); +static void mjIncludeXML(XMLElement* elem, string dir, const mjVFS* vfs, + std::unordered_set& included) { + // not an include, recursively go through all children + if (strcasecmp(elem->Value(), "include")) { + XMLElement* child = elem->FirstChildElement(); + for (; child; child = child->NextSiblingElement()) { + mjIncludeXML(child, dir, vfs, included); } - - // get filename - string filename; - mjXUtil::ReadAttrTxt(elem, "file", filename, true); - filename = dir + filename; - - // block repeated include files - for (size_t i=0; iFirstChildElement(); - if (!eleminc) { - throw mjXError(elem, "Empty include file '%s'", filename.c_str()); - } - - // get parent of - XMLElement* parent = (XMLElement*)elem->Parent(); - - // clone first child of included document, insert it after - XMLNode* first = parent->InsertAfterChild(elem, eleminc->DeepClone(parent->GetDocument())); - - // delete element, point to first - parent->DeleteChild(elem); - elem = first->ToElement(); - - // insert remaining elements from included document as siblings - eleminc = eleminc->NextSiblingElement(); - while (eleminc) { - elem = (XMLElement*)parent->InsertAfterChild(elem, eleminc->DeepClone(parent->GetDocument())); - eleminc = eleminc->NextSiblingElement(); - } - - // run XMLInclude on first new child - return mjIncludeXML(first->ToElement(), dir, vfs, included); + return; } - // otherwise check all child elements, return self - else { - XMLElement* child = elem->FirstChildElement(); - while (child) { - child = mjIncludeXML(child, dir, vfs, included); - if (child) { - child = child->NextSiblingElement(); - } + // make sure include has no children + if (!elem->NoChildren()) { + throw mjXError(elem, "Include element cannot have children"); + } + + // get filename + string filename; + mjXUtil::ReadAttrTxt(elem, "file", filename, true); + filename = dir + filename; + + // block repeated include files + if (included.find(filename) != included.end()) { + throw mjXError(elem, "File '%s' already included", filename.c_str()); + } + + // get data source + mjResource *resource = nullptr; + if ((resource = mju_openVfsResource(filename.c_str(), vfs)) == nullptr) { + // load from provider or OS filesystem + if ((resource = mju_openResource(filename.c_str())) == nullptr) { + throw mjXError(elem, "Could not open file '%s'", filename.c_str()); } - return elem; + } + + const char* xmlstring = nullptr; + int buffer_size = mju_readResource(resource, (const void**) &xmlstring); + if (buffer_size < 0) { + mju_closeResource(resource); + throw mjXError(elem, "Error reading file '%s'", filename.c_str()); + } else if (!buffer_size) { + mju_closeResource(resource); + throw mjXError(elem, "Empty file '%s'", filename.c_str()); + } + + // load XML file or parse string + XMLDocument doc; + doc.Parse(xmlstring, buffer_size); + + // close resource + mju_closeResource(resource); + + // check error + if (doc.Error()) { + char err[1000]; + mju::sprintf_arr(err, "XML parse error %d:\n%s\n", doc.ErrorID(), doc.ErrorStr()); + throw mjXError(elem, "Include error: '%s'", err); + } + + // remember that file was included + included.insert(filename); + + // get and check root element + XMLElement* docroot = doc.RootElement(); + if (!docroot) { + throw mjXError(elem, "Root element missing in file '%s'", filename.c_str()); + } + + // get and check first child + XMLElement* eleminc = docroot->FirstChildElement(); + if (!eleminc) { + throw mjXError(elem, "Empty include file '%s'", filename.c_str()); + } + + // get element + XMLElement* include = elem->ToElement(); + XMLDocument* include_doc = include->GetDocument(); + + // clone first child of included document + XMLNode* first = include->InsertFirstChild(eleminc->DeepClone(include_doc)); + + // point to first + XMLElement* child = first->ToElement(); + + // insert remaining elements from included document as siblings + eleminc = eleminc->NextSiblingElement(); + while (eleminc) { + child = include->InsertAfterChild(child, eleminc->DeepClone(include_doc))->ToElement(); + eleminc = eleminc->NextSiblingElement(); + } + + // recursively run include + child = include->FirstChildElement(); + for (; child; child = child->NextSiblingElement()) { + mjIncludeXML(child, dir, vfs, included); } } @@ -223,7 +220,7 @@ mjCModel* mjParseXML(const char* filename, const mjVFS* vfs, char* error, int er // check arguments if (!filename) { if (error) { - snprintf(error, error_sz, "mjParseXML: filename argument required\n"); + std::snprintf(error, error_sz, "mjParseXML: filename argument required\n"); } return nullptr; } @@ -241,7 +238,7 @@ mjCModel* mjParseXML(const char* filename, const mjVFS* vfs, char* error, int er // load from provider or fallback to OS filesystem if ((resource = mju_openResource(filename)) == nullptr) { if (error) { - snprintf(error, error_sz, "mjParseXML: could not open file '%s'", filename); + std::snprintf(error, error_sz, "mjParseXML: could not open file '%s'", filename); } return nullptr; } @@ -250,13 +247,13 @@ mjCModel* mjParseXML(const char* filename, const mjVFS* vfs, char* error, int er int buffer_size = mju_readResource(resource, (const void**) &xmlstring); if (buffer_size < 0) { if (error) { - snprintf(error, error_sz, "mjParseXML: error reading file '%s'", filename); + std::snprintf(error, error_sz, "mjParseXML: error reading file '%s'", filename); } mju_closeResource(resource); return nullptr; } else if (!buffer_size) { if (error) { - snprintf(error, error_sz, "mjParseXML: empty file '%s'", filename); + std::snprintf(error, error_sz, "mjParseXML: empty file '%s'", filename); } mju_closeResource(resource); return nullptr; @@ -303,8 +300,7 @@ mjCModel* mjParseXML(const char* filename, const mjVFS* vfs, char* error, int er try { if (!strcasecmp(root->Value(), "mujoco")) { // find include elements, replace them with subtree from xml file - vector included; - included.push_back(filename); + std::unordered_set included = {filename}; mjIncludeXML(root, model->modelfiledir, vfs, included); // parse MuJoCo model diff --git a/src/xml/xml_native_reader.cc b/src/xml/xml_native_reader.cc index 11c396ff..7b5f7476 100644 --- a/src/xml/xml_native_reader.cc +++ b/src/xml/xml_native_reader.cc @@ -27,9 +27,12 @@ #include #include -#include +#include "tinyxml2.h" + #include +#include #include +#include #include "engine/engine_plugin.h" #include "engine/engine_util_errmem.h" #include "engine/engine_util_misc.h" @@ -41,7 +44,6 @@ #include "user/user_util.h" #include "xml/xml_base.h" #include "xml/xml_util.h" -#include "tinyxml2.h" namespace { using std::string; @@ -50,7 +52,7 @@ using tinyxml2::XMLElement; void ReadPluginConfigs(tinyxml2::XMLElement* elem, mjCPlugin* pp) { std::map> config_attribs; - XMLElement* child = elem->FirstChildElement(); + XMLElement* child = FirstChildElement(elem); while (child) { std::string_view name = child->Value(); if (name == "config") { @@ -63,7 +65,7 @@ void ReadPluginConfigs(tinyxml2::XMLElement* elem, mjCPlugin* pp) { mjXUtil::ReadAttrTxt(child, "value", value, /* required = */ true); config_attribs[key] = value; } - child = child->NextSiblingElement(); + child = NextSiblingElement(child); } if (!pp && !config_attribs.empty()) { @@ -822,92 +824,92 @@ void mjXReader::Parse(XMLElement* root) { //------------------- parse MuJoCo sections embedded in all XML formats - for (XMLElement* section = root->FirstChildElement("compiler"); section; - section = section->NextSiblingElement("compiler")) { + for (XMLElement* section = FirstChildElement(root, "compiler"); section; + section = NextSiblingElement(section, "compiler")) { Compiler(section, model); } - for (XMLElement* section = root->FirstChildElement("option"); section; - section = section->NextSiblingElement("option")) { + for (XMLElement* section = FirstChildElement(root, "option"); section; + section = NextSiblingElement(section, "option")) { Option(section, &model->option); } - for (XMLElement* section = root->FirstChildElement("size"); section; - section = section->NextSiblingElement("size")) { + for (XMLElement* section = FirstChildElement(root, "size"); section; + section = NextSiblingElement(section, "size")) { Size(section, model); } //------------------ parse MJCF-specific sections - for (XMLElement* section = root->FirstChildElement("visual"); section; - section = section->NextSiblingElement("visual")) { + for (XMLElement* section = FirstChildElement(root, "visual"); section; + section = NextSiblingElement(section, "visual")) { Visual(section); } - for (XMLElement* section = root->FirstChildElement("statistic"); section; - section = section->NextSiblingElement("statistic")) { + for (XMLElement* section = FirstChildElement(root, "statistic"); section; + section = NextSiblingElement(section, "statistic")) { Statistic(section); } readingdefaults = true; - for (XMLElement* section = root->FirstChildElement("default"); section; - section = section->NextSiblingElement("default")) { + for (XMLElement* section = FirstChildElement(root, "default"); section; + section = NextSiblingElement(section, "default")) { Default(section, -1); } readingdefaults = false; - for (XMLElement* section = root->FirstChildElement("extension"); section; - section = section->NextSiblingElement("extension")) { + for (XMLElement* section = FirstChildElement(root, "extension"); section; + section = NextSiblingElement(section, "extension")) { Extension(section); } - for (XMLElement* section = root->FirstChildElement("custom"); section; - section = section->NextSiblingElement("custom")) { + for (XMLElement* section = FirstChildElement(root, "custom"); section; + section = NextSiblingElement(section, "custom")) { Custom(section); } - for (XMLElement* section = root->FirstChildElement("asset"); section; - section = section->NextSiblingElement("asset")) { + for (XMLElement* section = FirstChildElement(root, "asset"); section; + section = NextSiblingElement(section, "asset")) { Asset(section); } - for (XMLElement* section = root->FirstChildElement("worldbody"); section; - section = section->NextSiblingElement("worldbody")) { + for (XMLElement* section = FirstChildElement(root, "worldbody"); section; + section = NextSiblingElement(section, "worldbody")) { Body(section, &model->GetWorld()->spec, nullptr); } - for (XMLElement* section = root->FirstChildElement("contact"); section; - section = section->NextSiblingElement("contact")) { + for (XMLElement* section = FirstChildElement(root, "contact"); section; + section = NextSiblingElement(section, "contact")) { Contact(section); } - for (XMLElement* section = root->FirstChildElement("deformable"); section; - section = section->NextSiblingElement("deformable")) { + for (XMLElement* section = FirstChildElement(root, "deformable"); section; + section = NextSiblingElement(section, "deformable")) { Deformable(section); } - for (XMLElement* section = root->FirstChildElement("equality"); section; - section = section->NextSiblingElement("equality")) { + for (XMLElement* section = FirstChildElement(root, "equality"); section; + section = NextSiblingElement(section, "equality")) { Equality(section); } - for (XMLElement* section = root->FirstChildElement("tendon"); section; - section = section->NextSiblingElement("tendon")) { + for (XMLElement* section = FirstChildElement(root, "tendon"); section; + section = NextSiblingElement(section, "tendon")) { Tendon(section); } - for (XMLElement* section = root->FirstChildElement("actuator"); section; - section = section->NextSiblingElement("actuator")) { + for (XMLElement* section = FirstChildElement(root, "actuator"); section; + section = NextSiblingElement(section, "actuator")) { Actuator(section); } - for (XMLElement* section = root->FirstChildElement("sensor"); section; - section = section->NextSiblingElement("sensor")) { + for (XMLElement* section = FirstChildElement(root, "sensor"); section; + section = NextSiblingElement(section, "sensor")) { Sensor(section); } - for (XMLElement* section = root->FirstChildElement("keyframe"); section; - section = section->NextSiblingElement("keyframe")) { + for (XMLElement* section = FirstChildElement(root, "keyframe"); section; + section = NextSiblingElement(section, "keyframe")) { Keyframe(section); } } @@ -1044,7 +1046,7 @@ void mjXReader::Option(XMLElement* section, mjOption* opt) { text, false, false); for (int i=0; i < num_found; i++) { int group = disabled_act_groups[i]; - if (group < 0 ) { + if (group < 0) { throw mjXError(section, "disabled actuator group value must be non-negative"); } if (group > num_bitflags - 1) { @@ -1057,7 +1059,7 @@ void mjXReader::Option(XMLElement* section, mjOption* opt) { XMLElement* elem = FindSubElem(section, "flag"); if (elem) { #define READDSBL(NAME, MASK) \ - if( MapValue(elem, NAME, &n, enable_map, 2) ) { \ + if (MapValue(elem, NAME, &n, enable_map, 2)) { \ opt->disableflags ^= (opt->disableflags & MASK); \ opt->disableflags |= (n ? 0 : MASK); } @@ -1079,7 +1081,7 @@ void mjXReader::Option(XMLElement* section, mjOption* opt) { #undef READDSBL #define READENBL(NAME, MASK) \ - if( MapValue(elem, NAME, &n, enable_map, 2) ) { \ + if (MapValue(elem, NAME, &n, enable_map, 2)) { \ opt->enableflags ^= (opt->enableflags & MASK); \ opt->enableflags |= (n ? MASK : 0); } @@ -1303,7 +1305,7 @@ void mjXReader::OneFlex(XMLElement* elem, mjCFlex* pflex) { } // contact subelement - XMLElement* cont = elem->FirstChildElement("contact"); + XMLElement* cont = FirstChildElement(elem, "contact"); if (cont) { ReadAttrInt(cont, "contype", &pflex->contype); ReadAttrInt(cont, "conaffinity", &pflex->conaffinity); @@ -1323,7 +1325,7 @@ void mjXReader::OneFlex(XMLElement* elem, mjCFlex* pflex) { } // edge subelement - XMLElement* edge = elem->FirstChildElement("edge"); + XMLElement* edge = FirstChildElement(elem, "edge"); if (edge) { ReadAttr(edge, "stiffness", 1, &pflex->edgestiffness, text); ReadAttr(edge, "damping", 1, &pflex->edgedamping, text); @@ -1348,7 +1350,7 @@ void mjXReader::OneMesh(XMLElement* elem, mjCMesh* pmesh) { pmesh->set_refquat(ReadAttrArr(elem, "refquat")); pmesh->set_scale(ReadAttrArr(elem, "scale")); - XMLElement* eplugin = elem->FirstChildElement("plugin"); + XMLElement* eplugin = FirstChildElement(elem, "plugin"); if (eplugin) { OnePlugin(eplugin, &pmesh->plugin); } @@ -1400,7 +1402,7 @@ void mjXReader::OneSkin(XMLElement* elem, mjCSkin* pskin) { if (ReadAttrTxt(elem, "face", text)) String2Vector(text, pskin->face); // read bones - XMLElement* bone = elem->FirstChildElement("bone"); + XMLElement* bone = FirstChildElement(elem, "bone"); while (bone) { // read body ReadAttrTxt(bone, "body", text, true); @@ -1432,7 +1434,7 @@ void mjXReader::OneSkin(XMLElement* elem, mjCSkin* pskin) { pskin->vertweight.push_back(tempweight); // advance to next bone - bone = bone->NextSiblingElement("bone"); + bone = NextSiblingElement(bone, "bone"); } GetXMLPos(elem, pskin); @@ -1571,7 +1573,7 @@ void mjXReader::OneGeom(XMLElement* elem, mjmGeom* pgeom) { } // plugin sub-element - XMLElement* eplugin = elem->FirstChildElement("plugin"); + XMLElement* eplugin = FirstChildElement(elem, "plugin"); if (eplugin) { OnePlugin(eplugin, &pgeom->plugin); } @@ -2180,7 +2182,7 @@ void mjXReader::OneComposite(XMLElement* elem, mjmBody* pbody, mjCDef* def) { ReadAttr(elem, "flatinertia", 1, &comp.flatinertia, text); // plugin - XMLElement* eplugin = elem->FirstChildElement("plugin"); + XMLElement* eplugin = FirstChildElement(elem, "plugin"); if (eplugin) { ReadAttrTxt(eplugin, "plugin", comp.plugin_name); ReadAttrTxt(eplugin, "instance", comp.plugin_instance_name); @@ -2221,7 +2223,7 @@ void mjXReader::OneComposite(XMLElement* elem, mjmBody* pbody, mjCDef* def) { }; // skin - XMLElement* eskin = elem->FirstChildElement("skin"); + XMLElement* eskin = FirstChildElement(elem, "skin"); if (eskin) { comp.skin = true; if (MapValue(eskin, "texcoord", &n, bool_map, 2)) { @@ -2245,7 +2247,7 @@ void mjXReader::OneComposite(XMLElement* elem, mjmBody* pbody, mjCDef* def) { ReadAttr(elem, "solimpsmooth", mjNIMP, comp.solimpsmooth, text, false, false); // geom - XMLElement* egeom = elem->FirstChildElement("geom"); + XMLElement* egeom = FirstChildElement(elem, "geom"); if (egeom) { std::string material; mjmGeom& dgeom = comp.def[0].geom.spec; @@ -2273,7 +2275,7 @@ void mjXReader::OneComposite(XMLElement* elem, mjmBody* pbody, mjCDef* def) { } // site - XMLElement* esite = elem->FirstChildElement("site"); + XMLElement* esite = FirstChildElement(elem, "site"); if (esite) { std::string material; mjmSite& dsite = comp.def[0].site.spec; @@ -2285,7 +2287,7 @@ void mjXReader::OneComposite(XMLElement* elem, mjmBody* pbody, mjCDef* def) { } // joint - XMLElement* ejnt = elem->FirstChildElement("joint"); + XMLElement* ejnt = FirstChildElement(elem, "joint"); while (ejnt) { // kind int kind; @@ -2330,11 +2332,11 @@ void mjXReader::OneComposite(XMLElement* elem, mjmBody* pbody, mjCDef* def) { ReadAttr(ejnt, "frictionloss", 1, &el->joint.spec.frictionloss, text); // advance - ejnt = ejnt->NextSiblingElement("joint"); + ejnt = NextSiblingElement(ejnt, "joint"); } // tendon - XMLElement* eten = elem->FirstChildElement("tendon"); + XMLElement* eten = FirstChildElement(elem, "tendon"); while (eten) { // kind int kind; @@ -2366,11 +2368,11 @@ void mjXReader::OneComposite(XMLElement* elem, mjmBody* pbody, mjCDef* def) { ReadAttr(eten, "width", 1, &comp.def[kind].tendon.spec.width, text); // advance - eten = eten->NextSiblingElement("tendon"); + eten = NextSiblingElement(eten, "tendon"); } // pin - XMLElement* epin = elem->FirstChildElement("pin"); + XMLElement* epin = FirstChildElement(elem, "pin"); while (epin) { // read int coord[2] = {0, 0}; @@ -2381,7 +2383,7 @@ void mjXReader::OneComposite(XMLElement* elem, mjmBody* pbody, mjCDef* def) { comp.pin.push_back(coord[1]); // advance - epin = epin->NextSiblingElement("pin"); + epin = NextSiblingElement(epin, "pin"); } // make composite @@ -2444,7 +2446,7 @@ void mjXReader::OneFlexcomp(XMLElement* elem, mjmBody* pbody) { } // edge - XMLElement* edge = elem->FirstChildElement("edge"); + XMLElement* edge = FirstChildElement(elem, "edge"); if (edge) { if (MapValue(edge, "equality", &n, bool_map, 2)) { fcomp.equality = (n==1); @@ -2456,7 +2458,7 @@ void mjXReader::OneFlexcomp(XMLElement* elem, mjmBody* pbody) { } // contact - XMLElement* cont = elem->FirstChildElement("contact"); + XMLElement* cont = FirstChildElement(elem, "contact"); if (cont) { ReadAttrInt(cont, "contype", &fcomp.def.flex.contype); ReadAttrInt(cont, "conaffinity", &fcomp.def.flex.conaffinity); @@ -2476,7 +2478,7 @@ void mjXReader::OneFlexcomp(XMLElement* elem, mjmBody* pbody) { } // pin - XMLElement* epin = elem->FirstChildElement("pin"); + XMLElement* epin = FirstChildElement(elem, "pin"); while (epin) { // accumulate id, coord, range vector temp; @@ -2498,11 +2500,11 @@ void mjXReader::OneFlexcomp(XMLElement* elem, mjmBody* pbody) { } // advance - epin = epin->NextSiblingElement("pin"); + epin = NextSiblingElement(epin, "pin"); } // plugin - XMLElement* eplugin = elem->FirstChildElement("plugin"); + XMLElement* eplugin = FirstChildElement(elem, "plugin"); if (eplugin) { ReadAttrTxt(eplugin, "plugin", fcomp.plugin_name); ReadAttrTxt(eplugin, "instance", fcomp.plugin_instance_name); @@ -2579,7 +2581,7 @@ void mjXReader::Default(XMLElement* section, int parentid) { } // iterate over elements other than nested defaults - elem = section->FirstChildElement(); + elem = FirstChildElement(section); while (elem) { // get element name name = elem->Value(); @@ -2639,11 +2641,11 @@ void mjXReader::Default(XMLElement* section, int parentid) { mjm_finalize(def->tendon.spec.element); // advance - elem = elem->NextSiblingElement(); + elem = NextSiblingElement(elem); } // iterate over nested defaults - elem = section->FirstChildElement(); + elem = FirstChildElement(section); while (elem) { // get element name name = elem->Value(); @@ -2654,7 +2656,7 @@ void mjXReader::Default(XMLElement* section, int parentid) { } // advance - elem = elem->NextSiblingElement(); + elem = NextSiblingElement(elem); } } @@ -2662,7 +2664,7 @@ void mjXReader::Default(XMLElement* section, int parentid) { // extension section parser void mjXReader::Extension(XMLElement* section) { - XMLElement* elem = section->FirstChildElement(); + XMLElement* elem = FirstChildElement(section); while (elem) { // get sub-element name std::string_view name = elem->Value(); @@ -2687,7 +2689,7 @@ void mjXReader::Extension(XMLElement* section) { model->active_plugins.emplace_back(std::make_pair(plugin, plugin_slot)); } - XMLElement* child = elem->FirstChildElement(); + XMLElement* child = FirstChildElement(elem); while (child) { if (std::string(child->Value())=="instance") { if (model->hasImplicitPluginElem) { @@ -2704,12 +2706,12 @@ void mjXReader::Extension(XMLElement* section) { pp->plugin_slot = plugin_slot; pp->nstate = -1; // actual value to be filled in by the plugin later } - child = child->NextSiblingElement(); + child = NextSiblingElement(child); } } // advance to next element - elem = elem->NextSiblingElement(); + elem = NextSiblingElement(elem); } } @@ -2722,7 +2724,7 @@ void mjXReader::Custom(XMLElement* section) { double data[500]; // iterate over child elements - elem = section->FirstChildElement(); + elem = FirstChildElement(section); while (elem) { // get sub-element name name = elem->Value(); @@ -2784,7 +2786,7 @@ void mjXReader::Custom(XMLElement* section) { ReadAttrTxt(elem, "name", ptu->name, true); // read objects and add - XMLElement* obj = elem->FirstChildElement(); + XMLElement* obj = FirstChildElement(elem); while (obj) { // get sub-element name name = obj->Value(); @@ -2810,12 +2812,12 @@ void mjXReader::Custom(XMLElement* section) { } // advance to next object - obj = obj->NextSiblingElement(); + obj = NextSiblingElement(obj); } } // advance to next element - elem = elem->NextSiblingElement(); + elem = NextSiblingElement(elem); } } @@ -2828,7 +2830,7 @@ void mjXReader::Visual(XMLElement* section) { mjVisual* vis = &model->visual; // iterate over child elements - elem = section->FirstChildElement(); + elem = FirstChildElement(section); while (elem) { // get sub-element name name = elem->Value(); @@ -2946,7 +2948,7 @@ void mjXReader::Visual(XMLElement* section) { } // advance to next element - elem = elem->NextSiblingElement(); + elem = NextSiblingElement(elem); } } @@ -2959,7 +2961,7 @@ void mjXReader::Asset(XMLElement* section) { XMLElement* elem; // iterate over child elements - elem = section->FirstChildElement(); + elem = FirstChildElement(section); while (elem) { // get sub-element name name = elem->Value(); @@ -3091,7 +3093,7 @@ void mjXReader::Asset(XMLElement* section) { } // advance to next element - elem = elem->NextSiblingElement(); + elem = NextSiblingElement(elem); } } @@ -3114,7 +3116,7 @@ void mjXReader::Body(XMLElement* section, mjmBody* pbody, mjmFrame* frame) { } // iterate over sub-elements; attributes set while parsing parent body - elem = section->FirstChildElement(); + elem = FirstChildElement(section); while (elem) { // get sub-element name name = elem->Value(); @@ -3291,7 +3293,7 @@ void mjXReader::Body(XMLElement* section, mjmBody* pbody, mjmFrame* frame) { } // advance to next element - elem = elem->NextSiblingElement(); + elem = NextSiblingElement(elem); } } @@ -3303,7 +3305,7 @@ void mjXReader::Contact(XMLElement* section) { XMLElement* elem; // iterate over child elements - elem = section->FirstChildElement(); + elem = FirstChildElement(section); while (elem) { // get sub-element name name = elem->Value(); @@ -3333,7 +3335,7 @@ void mjXReader::Contact(XMLElement* section) { } // advance to next element - elem = elem->NextSiblingElement(); + elem = NextSiblingElement(elem); } } @@ -3344,7 +3346,7 @@ void mjXReader::Equality(XMLElement* section) { XMLElement* elem; // iterate over child elements - elem = section->FirstChildElement(); + elem = FirstChildElement(section); while (elem) { // get class if specified, otherwise use default0 mjCDef* def = GetClass(elem); @@ -3357,7 +3359,7 @@ void mjXReader::Equality(XMLElement* section) { OneEquality(elem, pequality); // advance to next element - elem = elem->NextSiblingElement(); + elem = NextSiblingElement(elem); } } @@ -3369,7 +3371,7 @@ void mjXReader::Deformable(XMLElement* section) { XMLElement* elem; // iterate over child elements - elem = section->FirstChildElement(); + elem = FirstChildElement(section); while (elem) { // get sub-element name name = elem->Value(); @@ -3395,7 +3397,7 @@ void mjXReader::Deformable(XMLElement* section) { } // advance to next element - elem = elem->NextSiblingElement(); + elem = NextSiblingElement(elem); } } @@ -3408,7 +3410,7 @@ void mjXReader::Tendon(XMLElement* section) { double data; // iterate over child elements - elem = section->FirstChildElement(); + elem = FirstChildElement(section); while (elem) { // get class if specified, otherwise use default0 mjCDef* def = GetClass(elem); @@ -3421,7 +3423,7 @@ void mjXReader::Tendon(XMLElement* section) { OneTendon(elem, pten); // process wrap sub-elements - XMLElement* sub = elem->FirstChildElement(); + XMLElement* sub = FirstChildElement(elem); while (sub) { // get wrap type string wrap = sub->Value(); @@ -3459,11 +3461,11 @@ void mjXReader::Tendon(XMLElement* section) { mjm_setString(pwrap->info, ("line = " + std::to_string(sub->GetLineNum())).c_str()); // advance to next sub-element - sub = sub->NextSiblingElement(); + sub = NextSiblingElement(sub); } // advance to next element - elem = elem->NextSiblingElement(); + elem = NextSiblingElement(elem); } } @@ -3474,7 +3476,7 @@ void mjXReader::Actuator(XMLElement* section) { XMLElement* elem; // iterate over child elements - elem = section->FirstChildElement(); + elem = FirstChildElement(section); while (elem) { // get class if specified, otherwise use default0 mjCDef* def = GetClass(elem); @@ -3487,7 +3489,7 @@ void mjXReader::Actuator(XMLElement* section) { OneActuator(elem, pact); // advance to next element - elem = elem->NextSiblingElement(); + elem = NextSiblingElement(elem); } } @@ -3496,7 +3498,7 @@ void mjXReader::Actuator(XMLElement* section) { // sensor section parser void mjXReader::Sensor(XMLElement* section) { int n; - XMLElement* elem = section->FirstChildElement(); + XMLElement* elem = FirstChildElement(section); while (elem) { // create sensor, get string type mjmSensor* psen = mjm_addSensor(model); @@ -3799,7 +3801,7 @@ void mjXReader::Sensor(XMLElement* section) { std::string("line = " + std::to_string(elem->GetLineNum()) + ", column = -1").c_str()); // advance to next element - elem = elem->NextSiblingElement(); + elem = NextSiblingElement(elem); } } @@ -3813,7 +3815,7 @@ void mjXReader::Keyframe(XMLElement* section) { double data[1000]; // iterate over child elements - elem = section->FirstChildElement(); + elem = FirstChildElement(section); while (elem) { // add keyframe mjCKey* pk = model->AddKey(); @@ -3865,7 +3867,7 @@ void mjXReader::Keyframe(XMLElement* section) { } // advance to next element - elem = elem->NextSiblingElement(); + elem = NextSiblingElement(elem); } } diff --git a/src/xml/xml_util.cc b/src/xml/xml_util.cc index c838d8d6..a58487d6 100644 --- a/src/xml/xml_util.cc +++ b/src/xml/xml_util.cc @@ -13,7 +13,6 @@ // limitations under the License. #include -#include #include #include #include @@ -24,11 +23,15 @@ #include #include #include +#include #include #include #include +#include #include +#include "tinyxml2.h" + #include "cc/array_safety.h" #include "engine/engine_util_errmem.h" #include "xml/xml_util.h" @@ -111,23 +114,52 @@ mjXError::mjXError(const XMLElement* elem, const char* msg, const char* str, int //---------------------------------- class mjXSchema implementation -------------------------------- -// constructor -mjXSchema::mjXSchema(const char* schema[][mjXATTRNUM], int nrow, bool checkptr) { - // clear fields - name.clear(); - type = '?'; - child.clear(); - attr.clear(); - error.clear(); +XMLElement* FirstChildElement(XMLElement* e, const char* name) { + XMLElement* child = e->FirstChildElement(); + for (; child; child = child->NextSiblingElement()) { + if (!std::strcmp(child->Name(), "include")) { + XMLElement* temp = FirstChildElement(child, name); + if (temp) { + return temp; + } + continue; + } - // checks nrow and first element - if (nrow<1) { - error = "number of rows must be positive"; - return; + if (!name || !std::strcmp(child->Name(), name)) { + return child; + } } - if (schema[0][0][0]=='<' || schema[0][0][0]=='>') { - error = "expected element, found bracket"; - return; + return nullptr; +} + +XMLElement* NextSiblingElement(XMLElement* e, const char* name) { + XMLElement* elem = e->NextSiblingElement(); + for (; elem; elem = elem->NextSiblingElement()) { + if (!std::strcmp(elem->Name(), "include")) { + XMLElement* temp = FirstChildElement(elem, name); + if (temp) { + return temp; + } + continue; + } + + if (!name || !std::strcmp(elem->Name(), name)) { + return elem; + } + } + + XMLElement* parent = e->Parent()->ToElement(); + if (parent && !std::strcmp(parent->Name(), "include")) { + return NextSiblingElement(parent, name); + } + + return nullptr; +} + +// constructor +mjXSchema::mjXSchema(const char* schema[][mjXATTRNUM], unsigned nrow, bool checkptr) { + if (schema[0][0][0] == '<' || schema[0][0][0] == '>') { + throw "expected element, found bracket"; } // check entire schema for null pointers @@ -138,8 +170,7 @@ mjXSchema::mjXSchema(const char* schema[][mjXATTRNUM], int nrow, bool checkptr) // base pointers if (!schema[i][0]) { mju::sprintf_arr(msg, "null pointer found in row %d", i); - error = msg; - return; + throw msg; } // detect element @@ -147,33 +178,29 @@ mjXSchema::mjXSchema(const char* schema[][mjXATTRNUM], int nrow, bool checkptr) // first 3 pointers required if (!schema[i][1] || !schema[i][2]) { mju::sprintf_arr(msg, "null pointer in row %d, element %s", i, schema[i][0]); - error = msg; - return; + throw msg; } // check type if (schema[i][1][0]!='!' && schema[i][1][0]!='?' && schema[i][1][0]!='*' && schema[i][1][0]!='R') { mju::sprintf_arr(msg, "invalid type in row %d, element %s", i, schema[i][0]); - error = msg; - return; + throw msg; } // number of attributes int nattr = atoi(schema[i][2]); - if (nattr<0 || nattr>mjXATTRNUM-3) { + if (nattr < 0 || nattr > mjXATTRNUM-3) { mju::sprintf_arr(msg, "invalid number of attributes in row %d, element %s", i, schema[i][0]); - error = msg; - return; + throw msg; } // attribute pointers for (int j=0; j1) { // check for bracketed block if (schema[1][0][0]!='<' || schema[nrow-1][0][0]!='>') { - error = "expected brackets after complex element"; - return; + throw "expected brackets after complex element"; } // parse block into simple and complex elements, create children @@ -222,18 +248,12 @@ mjXSchema::mjXSchema(const char* schema[][mjXATTRNUM], int nrow, bool checkptr) // closing bracket not found if (end > nrow-1) { - error = "matching closing bracket not found"; - return; + throw "matching closing bracket not found"; } } // add element, check for error - mjXSchema* elem = new mjXSchema(schema+start, end-start+1, false); - child.push_back(elem); - if (!elem->error.empty()) { - error = elem->error; - return; - } + subschema_.emplace_back(schema+start, end-start+1, false); // proceed with next subelement start = end+1; @@ -243,23 +263,8 @@ mjXSchema::mjXSchema(const char* schema[][mjXATTRNUM], int nrow, bool checkptr) -// destructor -mjXSchema::~mjXSchema() { - // delete children recursively - for (unsigned int i=0; i60) { str << "\n"; printspace(str, (cnt = std::max(30, baselen)), " "); - } - str << attr[i] << " "; - cnt += (int)attr[i].size() + 1; + str << attr << " "; + cnt += (int)attr.size() + 1; } str << "\n"; // children - for (int i=0; i<(int)child.size(); i++) { - child[i]->Print(str, level+1); + for (const mjXSchema& subschema : subschema_) { + subschema.Print(str, level+1); } } // print schema as HTML table -void mjXSchema::PrintHTML(std::stringstream& str, int level, bool pad) { +void mjXSchema::PrintHTML(std::stringstream& str, int level, bool pad) const { // replace body with (world)body - string name1 = (name=="body" ? "(world)body" : name); + string name1 = (name_ == "body" ? "(world)body" : name_); // open table if (level==0) { @@ -335,13 +339,13 @@ void mjXSchema::PrintHTML(std::stringstream& str, int level, bool pad) { } // type - str << "\t" << type << "\n"; + str << "\t" << type_ << "\n"; // attributes str << "\t"; - if (!attr.empty()) { - for (int i=0; i<(int)attr.size(); i++) { - str << attr[i] << " "; + if (!attr_.empty()) { + for (const std::string& attr : attr_) { + str << attr << " "; } } else { str << "no attributes"; @@ -349,12 +353,12 @@ void mjXSchema::PrintHTML(std::stringstream& str, int level, bool pad) { str << "\n\n"; // children - for (int i=0; i<(int)child.size(); i++) { - child[i]->PrintHTML(str, level+1, pad); + for (const mjXSchema& subschema : subschema_) { + subschema.PrintHTML(str, level+1, pad); } // close table - if (level==0) { + if (!level) { str << "\n"; } } @@ -363,23 +367,16 @@ void mjXSchema::PrintHTML(std::stringstream& str, int level, bool pad) { // check for name match bool mjXSchema::NameMatch(XMLElement* elem, int level) { - // special handling of body and worldbody - if (name=="body") { - if (level==1 && !strcmp(elem->Value(), "worldbody")) { - return true; - } - - if (level!=1 && !strcmp(elem->Value(), "body")) { - return true; - } - - if (level>=1 && !strcmp(elem->Value(), "frame")) { - return true; - } + // special handling of body, worldbody, and frame + if (name_ == "body" && + ((level == 1 && !strcmp(elem->Value(), "worldbody")) || + (level != 1 && !strcmp(elem->Value(), "body")) || + (level >= 1 && !strcmp(elem->Value(), "frame")))) { + return true; } // regular check - return (name==elem->Value()); + return name_ == elem->Value(); } @@ -403,87 +400,73 @@ XMLElement* mjXSchema::Check(XMLElement* elem, int level) { // check attributes const XMLAttribute* attribute = elem->FirstAttribute(); - while (attribute) { - missing = true; - for (int i=0; i<(int)attr.size(); i++) { - if (attr[i]==attribute->Name()) { - missing = false; - break; - } - } - if (missing) { + for (; attribute != nullptr; attribute = attribute->Next()) { + if (attr_.find(attribute->Name()) == attr_.end()) { error = "unrecognized attribute: '" + string(attribute->Name()) + "'"; return elem; } - - // next attribute - attribute = attribute->Next(); } // handle recursion - if (type=='R') { - // loop over sub-elements with same name - sub = elem->FirstChildElement((const char*)name.c_str()); - while (sub) { - // check sub-tree + if (type_ == 'R') { + // check child elements with same name + sub = FirstChildElement(elem, name_.c_str()); + for (; sub != nullptr; sub = NextSiblingElement(sub, name_.c_str())) { if ((bad = Check(sub, level+1))) { return bad; } - - // advance to next sub-element with same name - sub = sub->NextSiblingElement((const char*)name.c_str()); } } // clear reference counts - for (int i=0; i<(int)child.size(); i++) { - child[i]->refcnt = 0; + for (mjXSchema& subschema : subschema_) { + subschema.refcnt_ = 0; } // check sub-elements, update refcnt - sub = elem->FirstChildElement(); - while (sub) { - // find in child array, update refcnt + sub = FirstChildElement(elem); + for (; sub != nullptr; sub = NextSiblingElement(sub)) { missing = true; - for (int i=0; i<(int)child.size(); i++) { - if (child[i]->NameMatch(sub, level+1)) { + + for (mjXSchema& subschema : subschema_) { + if (subschema.NameMatch(sub, level+1)) { // check sub-tree - if ((bad = child[i]->Check(sub, level+1))) { - error = child[i]->error; + if ((bad = subschema.Check(sub, level+1))) { + error = subschema.error; return bad; } // mark found missing = false; - child[i]->refcnt++; + subschema.refcnt_++; break; } } // missing, unless recursive - if (missing && !(type=='R' && NameMatch(sub, level+1))) { + if (missing && !(type_ == 'R' && NameMatch(sub, level+1))) { error = "unrecognized element"; return sub; } - - // advance to next sub-element - sub = sub->NextSiblingElement(); } // enforce sub-element types - msg[0] = 0; - for (int i=0; i<(int)child.size(); i++) { - switch (child[i]->type) { + msg[0] = '\0'; + for (mjXSchema& subschema : subschema_) { + switch (subschema.type_) { case '!': - if (child[i]->refcnt != 1) - mju::sprintf_arr(msg, "required sub-element '%s' found %d time(s)", - child[i]->name.c_str(), child[i]->refcnt); + if (subschema.refcnt_ > 1) + mju::sprintf_arr(msg, "unique element '%s' found %d times", + subschema.name_.c_str(), subschema.refcnt_); + else if (subschema.refcnt_ < 1) + mju::sprintf_arr(msg, "element '%s' is required", + subschema.name_.c_str()); break; case '?': - if (child[i]->refcnt > 1) - mju::sprintf_arr(msg, "unique sub-element '%s' found %d time(s)", - child[i]->name.c_str(), child[i]->refcnt); + if (subschema.refcnt_ > 1) + mju::sprintf_arr(msg, "unique element '%s' found %d times", + subschema.name_.c_str(), subschema.refcnt_); break; default: @@ -495,9 +478,8 @@ XMLElement* mjXSchema::Check(XMLElement* elem, int level) { if (msg[0]) { error = msg; return elem; - } else { - return 0; } + return nullptr; } @@ -557,7 +539,7 @@ template bool mjXUtil::ReadAttrValues(XMLElement* elem, const char* attr, template bool mjXUtil::ReadAttrValues(XMLElement* elem, const char* attr, std::function push, int max); template bool mjXUtil::ReadAttrValues(XMLElement* elem, const char* attr, - std::function push, int max); + std::function push, int max); @@ -581,7 +563,7 @@ bool mjXUtil::SameVector(const T* vec1, const T* vec2, int n) { template bool mjXUtil::SameVector(const double* vec1, const double* vec2, int n); template bool mjXUtil::SameVector(const float* vec1, const float* vec2, int n); template bool mjXUtil::SameVector(const int* vec1, const int* vec2, int n); -template bool mjXUtil::SameVector(const mjtByte* vec1, const mjtByte* vec2, int n); +template bool mjXUtil::SameVector(const unsigned char* vec1, const unsigned char* vec2, int n); // find string in map, return corresponding integer (-1: not found) @@ -633,7 +615,7 @@ template std::optional> mjXUtil::ReadAttrVec(XMLElement* elem, const char* attr, bool required); template std::optional> mjXUtil::ReadAttrVec(XMLElement* elem, const char* attr, bool required); -template std::optional> +template std::optional> mjXUtil::ReadAttrVec(XMLElement* elem, const char* attr, bool required); @@ -675,7 +657,7 @@ template std::optional mjXUtil::ReadAttrNum(XMLElement* elem, const char* attr, bool required); template std::optional mjXUtil::ReadAttrNum(XMLElement* elem, const char* attr, bool required); -template std::optional +template std::optional mjXUtil::ReadAttrNum(XMLElement* elem, const char* attr, bool required); @@ -706,17 +688,18 @@ int mjXUtil::ReadAttr(XMLElement* elem, const char* attr, const int len, return maybe_vec->size(); } -template int mjXUtil::ReadAttr(XMLElement* elem, const char* attr, const int len, +template int mjXUtil::ReadAttr(XMLElement* elem, const char* attr, int len, double* data, string& text, bool required, bool exact); -template int mjXUtil::ReadAttr(XMLElement* elem, const char* attr, const int len, +template int mjXUtil::ReadAttr(XMLElement* elem, const char* attr, int len, float* data, string& text, bool required, bool exact); -template int mjXUtil::ReadAttr(XMLElement* elem, const char* attr, const int len, +template int mjXUtil::ReadAttr(XMLElement* elem, const char* attr, int len, int* data, string& text, bool required, bool exact); -template int mjXUtil::ReadAttr(XMLElement* elem, const char* attr, const int len, - mjtByte* data, string& text, bool required, bool exact); +template int mjXUtil::ReadAttr(XMLElement* elem, const char* attr, int len, + unsigned char* data, string& text, bool required, + bool exact); // read quaternion attribute // throw error if identically zero @@ -725,7 +708,7 @@ int mjXUtil::ReadQuat(XMLElement* elem, const char* attr, double* data, string& ReadAttr(elem, attr, /*len=*/4, data, text, required, /*exact=*/true); // check for 0 quaternion - if (data[0] == 0 && data[1] == 0 && data[2] == 0 && data[3] == 0 ) { + if (data[0] == 0 && data[1] == 0 && data[2] == 0 && data[3] == 0) { throw mjXError(elem, "zero quaternion is not allowed"); } @@ -1030,7 +1013,8 @@ template void mjXUtil::WriteAttr(XMLElement* elem, string name, int n, const int* data, const int* def); template void mjXUtil::WriteAttr(XMLElement* elem, string name, int n, - const mjtByte* data, const mjtByte* def); + const unsigned char* data, + const unsigned char* def); // write vector attribute, default = zero array diff --git a/src/xml/xml_util.h b/src/xml/xml_util.h index dd0f6b78..43f71290 100644 --- a/src/xml/xml_util.h +++ b/src/xml/xml_util.h @@ -19,20 +19,22 @@ #include #include #include +#include #include #include #include -#include - - -// TinyXML #include "tinyxml2.h" // error string copy void mjCopyError(char* dst, const char* src, int maxlen); +using tinyxml2::XMLElement; + +XMLElement* FirstChildElement(XMLElement* e, const char* name = nullptr); +XMLElement* NextSiblingElement(XMLElement* e, const char* name = nullptr); + // XML Error info class [[nodiscard]] mjXError { @@ -54,24 +56,22 @@ class [[nodiscard]] mjXError { // Custom XML file validation class mjXSchema { public: - mjXSchema(const char* schema[][mjXATTRNUM], // constructor - int nrow, bool checkptr = true); - ~mjXSchema(); // destructor + mjXSchema(const char* schema[][mjXATTRNUM], unsigned nrow, bool checkptr = true); - std::string GetError(void); // return error - void Print(std::stringstream& str, int level); // print schema - void PrintHTML(std::stringstream& str, int level, bool pad); + std::string GetError(); // return error + void Print(std::stringstream& str, int level) const; // print schema + void PrintHTML(std::stringstream& str, int level, bool pad) const; bool NameMatch(tinyxml2::XMLElement* elem, int level); // does name match tinyxml2::XMLElement* Check(tinyxml2::XMLElement* elem, int level); // validator private: - std::string name; // element name - char type; // element type: '?', '!', '*', 'R' - std::vector attr; // allowed attributes - std::vector child; // allowed child elements + std::string name_; // element name + char type_; // element type: '?', '!', '*', 'R' + std::set attr_; // allowed attributes + std::vector subschema_; // allowed child elements - int refcnt; // refcount used for validation + int refcnt_ = 0; // refcount used for validation std::string error; // error from constructor or Check }; @@ -141,7 +141,7 @@ class mjXUtil { // deprecated: use ReadAttrVec or ReadAttrArr template - static int ReadAttr(tinyxml2::XMLElement* elem, const char* attr, const int len, + static int ReadAttr(tinyxml2::XMLElement* elem, const char* attr, int len, T* data, std::string& text, bool required = false, bool exact = true); diff --git a/test/xml/xml_native_reader_test.cc b/test/xml/xml_native_reader_test.cc index 9e51bf24..b85aa8db 100644 --- a/test/xml/xml_native_reader_test.cc +++ b/test/xml/xml_native_reader_test.cc @@ -15,7 +15,9 @@ // Tests for xml/xml_native_reader.cc. #include +#include #include +#include #include #include @@ -40,6 +42,22 @@ using ::testing::FloatEq; using XMLReaderTest = MujocoTest; +TEST_F(XMLReaderTest, UniqueElementTest) { + std::array error; + static constexpr char xml[] = R"( + + + + )"; + + mjModel* model = LoadModelFromString(xml, error.data(), error.size()); + ASSERT_THAT(model, IsNull()); + EXPECT_THAT(error.data(), HasSubstr("unique element 'flag' found 2 times")); +} + TEST_F(XMLReaderTest, MemorySize) { std::array error; { @@ -419,8 +437,103 @@ TEST_F(XMLReaderTest, InvalidDoubleOrientation) { } } } +// ------------------------ test including ------------------------------------- -// ---------------------- test frame parsing --------------------------------- +TEST_F(XMLReaderTest, IncludeTest) { + static constexpr char xml[] = R"( + + + + + + + )"; + + static constexpr char xml1[] = R"( + + + )"; + + static constexpr char xml2[]= R"( + + + + )"; + + static constexpr char xml3[]= R"( + + + )"; + + auto vfs = std::make_unique(); + mj_defaultVFS(vfs.get()); + + mj_makeEmptyFileVFS(vfs.get(), "model1.xml", sizeof(xml1)); + std::memcpy(vfs->filedata[vfs->nfile - 1], xml1, sizeof(xml1)); + + mj_makeEmptyFileVFS(vfs.get(), "model2.xml", sizeof(xml2)); + std::memcpy(vfs->filedata[vfs->nfile - 1], xml2, sizeof(xml2)); + + mj_makeEmptyFileVFS(vfs.get(), "model3.xml", sizeof(xml3)); + std::memcpy(vfs->filedata[vfs->nfile - 1], xml3, sizeof(xml3)); + + std::array error; + mjModel* model = LoadModelFromString(xml, error.data(), + error.size(), vfs.get()); + ASSERT_THAT(model, NotNull()); + EXPECT_EQ(mj_name2id(model, mjOBJ_GEOM, "ball"), 2); + EXPECT_EQ(mj_name2id(model, mjOBJ_GEOM, "another_box"), 3); + mj_deleteModel(model); + mj_deleteVFS(vfs.get()); +} + +TEST_F(XMLReaderTest, IncludeChildTest) { + static constexpr char xml[] = R"( + + + + + + + + )"; + + std::array error; + mjModel* model = LoadModelFromString(xml, error.data(), error.size()); + ASSERT_THAT(model, IsNull()); + EXPECT_THAT(error.data(), HasSubstr("Include element cannot have children")); + mj_deleteModel(model); +} + +TEST_F(XMLReaderTest, IncludeSameFileTest) { + static constexpr char xml[] = R"( + + + + )"; + + static constexpr char xml1[] = R"( + + + )"; + + + auto vfs = std::make_unique(); + mj_defaultVFS(vfs.get()); + + mj_makeEmptyFileVFS(vfs.get(), "model1.xml", sizeof(xml1)); + std::memcpy(vfs->filedata[vfs->nfile - 1], xml1, sizeof(xml1)); + + std::array error; + mjModel* model = LoadModelFromString(xml, error.data(), error.size(), + vfs.get()); + ASSERT_THAT(model, IsNull()); + EXPECT_THAT(error.data(), HasSubstr("File 'model1.xml' already included")); + mj_deleteModel(model); + mj_deleteVFS(vfs.get()); +} + +// ------------------------ test frame parsing --------------------------------- TEST_F(XMLReaderTest, ParseFrame) { static constexpr char xml[] = R"( @@ -453,7 +566,7 @@ TEST_F(XMLReaderTest, ParseFrame) { mj_deleteModel(m); } -// ---------------------- test camera parsing --------------------------------- +// ----------------------- test camera parsing --------------------------------- TEST_F(XMLReaderTest, CameraInvalidFovyAndSensorsize) { static constexpr char xml[] = R"( @@ -509,7 +622,7 @@ TEST_F(XMLReaderTest, CameraSensorsizeRequiresResolution) { EXPECT_THAT(error.data(), HasSubstr("line 6")); } -// ---------------------- test inertia parsing -------------------------------- +// ----------------------- test inertia parsing -------------------------------- TEST_F(XMLReaderTest, InvalidInertialOrientation) { static constexpr char xml[] = R"(