Keep include elements while parsing XML. No change in behavior.

PiperOrigin-RevId: 606203141
Change-Id: Ica06f7c121a90597cd29f74692230365edd4302f
This commit is contained in:
Kyle Bayes
2024-02-12 03:53:56 -08:00
committed by Copybara-Service
parent 4e5b194e88
commit ff39d0b812
5 changed files with 460 additions and 365 deletions
+101 -105
View File
@@ -20,25 +20,27 @@
#include <xlocale.h>
#endif
#include <cstdio>
#include <string>
#include <unordered_set>
#include "tinyxml2.h"
#include <mujoco/mjmodel.h>
#include <mujoco/mjplugin.h>
#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<string>& 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<string>& 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; i<included.size(); i++) {
if (!strcasecmp(included[i].c_str(), filename.c_str())) {
throw mjXError(elem, "File '%s' already included", filename.c_str());
}
}
// get data source
mjResource *resource = nullptr;
const char* xmlstring = 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());
}
}
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.push_back(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 parent of <include>
XMLElement* parent = (XMLElement*)elem->Parent();
// clone first child of included document, insert it after <include>
XMLNode* first = parent->InsertAfterChild(elem, eleminc->DeepClone(parent->GetDocument()));
// delete <include> 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 <include> 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<string> included;
included.push_back(filename);
std::unordered_set<string> included = {filename};
mjIncludeXML(root, model->modelfiledir, vfs, included);
// parse MuJoCo model
+98 -96
View File
@@ -27,9 +27,12 @@
#include <utility>
#include <vector>
#include <mujoco/mjmacro.h>
#include "tinyxml2.h"
#include <mujoco/mjmodel.h>
#include <mujoco/mjplugin.h>
#include <mujoco/mjvisualize.h>
#include <mujoco/mjtnum.h>
#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<std::string, std::string, std::less<>> 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<double, 4>(elem, "refquat"));
pmesh->set_scale(ReadAttrArr<double, 3>(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<int> 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);
}
}
+129 -145
View File
@@ -13,7 +13,6 @@
// limitations under the License.
#include <algorithm>
#include <array>
#include <climits>
#include <cmath>
#include <cstddef>
@@ -24,11 +23,15 @@
#include <iostream>
#include <limits>
#include <optional>
#include <set>
#include <sstream>
#include <string>
#include <type_traits>
#include <utility>
#include <vector>
#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; j<nattr; j++) {
if (!schema[i][3+j]) {
mju::sprintf_arr(msg, "null attribute %d in row %d, element %s", j, i, schema[i][0]);
error = msg;
return;
throw msg;
}
}
}
@@ -181,21 +208,20 @@ mjXSchema::mjXSchema(const char* schema[][mjXATTRNUM], int nrow, bool checkptr)
}
// set name and type
name = schema[0][0];
type = schema[0][1][0];
name_ = schema[0][0];
type_ = schema[0][1][0];
// set attributes
int nattr = atoi(schema[0][2]);
for (int i=0; i<nattr; i++) {
attr.push_back(schema[0][3+i]);
for (int i = 0; i < nattr; i++) {
attr_.emplace(schema[0][3 + i]);
}
// process sub-elements of complex element
if (nrow>1) {
// 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; i<child.size(); i++) {
delete child[i];
}
// clear fields
child.clear();
attr.clear();
error.clear();
}
// get pointer to error message
string mjXSchema::GetError(void) {
string mjXSchema::GetError() {
return error;
}
@@ -275,13 +280,13 @@ static void printspace(std::stringstream& str, int n, const char* space) {
// print schema as text
void mjXSchema::Print(std::stringstream& str, int level) {
void mjXSchema::Print(std::stringstream& str, int level) const {
// replace body with (world)body
string name1 = (name=="body" ? "(world)body" : name);
string name1 = (name_ == "body") ? "(world)body" : name_;
// space, name, type
printspace(str, 3*level, " ");
str << name1 << " (" << type << ")";
str << name1 << " (" << type_ << ")";
int baselen = 3*level + (int)name1.size() + 4;
if (baselen<30) {
printspace(str, 30-baselen, " ");
@@ -289,30 +294,29 @@ void mjXSchema::Print(std::stringstream& str, int level) {
// attributes
int cnt = std::max(baselen, 30);
for (int i=0; i<(int)attr.size(); i++) {
for (const std::string& attr : attr_) {
if (cnt>60) {
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<td class=\"ty\">" << type << "</td>\n";
str << "\t<td class=\"ty\">" << type_ << "</td>\n";
// attributes
str << "\t<td class=\"at\">";
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 << "<span style=\"color:black\"><i>no attributes</i></span>";
@@ -349,12 +353,12 @@ void mjXSchema::PrintHTML(std::stringstream& str, int level, bool pad) {
str << "</td>\n</tr>\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 << "</table>\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<void (int, int)> push, int max);
template bool mjXUtil::ReadAttrValues(XMLElement* elem, const char* attr,
std::function<void (int, mjtByte)> push, int max);
std::function<void (int, unsigned char)> 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<std::vector<float>>
mjXUtil::ReadAttrVec(XMLElement* elem, const char* attr, bool required);
template std::optional<std::vector<int>>
mjXUtil::ReadAttrVec(XMLElement* elem, const char* attr, bool required);
template std::optional<std::vector<mjtByte>>
template std::optional<std::vector<unsigned char>>
mjXUtil::ReadAttrVec(XMLElement* elem, const char* attr, bool required);
@@ -675,7 +657,7 @@ template std::optional<float>
mjXUtil::ReadAttrNum(XMLElement* elem, const char* attr, bool required);
template std::optional<int>
mjXUtil::ReadAttrNum(XMLElement* elem, const char* attr, bool required);
template std::optional<mjtByte>
template std::optional<unsigned char>
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<double> attribute, default = zero array
+16 -16
View File
@@ -19,20 +19,22 @@
#include <array>
#include <functional>
#include <optional>
#include <set>
#include <string>
#include <vector>
#include <sstream>
#include <mujoco/mjmodel.h>
// 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<std::string> attr; // allowed attributes
std::vector<mjXSchema*> child; // allowed child elements
std::string name_; // element name
char type_; // element type: '?', '!', '*', 'R'
std::set<std::string> attr_; // allowed attributes
std::vector<mjXSchema> 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<typename T>
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);
+116 -3
View File
@@ -15,7 +15,9 @@
// Tests for xml/xml_native_reader.cc.
#include <array>
#include <cstring>
#include <limits>
#include <memory>
#include <string>
#include <vector>
@@ -40,6 +42,22 @@ using ::testing::FloatEq;
using XMLReaderTest = MujocoTest;
TEST_F(XMLReaderTest, UniqueElementTest) {
std::array<char, 1024> error;
static constexpr char xml[] = R"(
<mujoco>
<option>
<flag sensor="disable"/>
<flag sensor="disable"/>
</option>
</mujoco>
)";
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<char, 1024> error;
{
@@ -419,8 +437,103 @@ TEST_F(XMLReaderTest, InvalidDoubleOrientation) {
}
}
}
// ------------------------ test including -------------------------------------
// ---------------------- test frame parsing ---------------------------------
TEST_F(XMLReaderTest, IncludeTest) {
static constexpr char xml[] = R"(
<mujoco>
<worldbody>
<geom name="plane" type="plane" size="1 1 1"/>
<include file="model1.xml"/>
<include file="model2.xml"/>
</worldbody>
</mujoco>)";
static constexpr char xml1[] = R"(
<mujoco>
<geom name="box" type="box" size="1 1 1"/>
</mujoco>)";
static constexpr char xml2[]= R"(
<mujoco>
<geom name="ball" type="sphere" size="2"/>
<include file="model3.xml"/>
</mujoco>)";
static constexpr char xml3[]= R"(
<mujoco>
<geom name="another_box" type="box" size="2 2 2"/>
</mujoco>)";
auto vfs = std::make_unique<mjVFS>();
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<char, 1024> 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"(
<mujoco>
<worldbody>
<geom name="plane" type="plane" size="1 1 1"/>
<include file="model1.xml">
<geom name="box" type="box" size="1 1 1"/>
</include>
</worldbody>
</mujoco>)";
std::array<char, 1024> 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"(
<mujoco>
<include file="model1.xml"/>
<include file="model1.xml"/>
</mujoco>)";
static constexpr char xml1[] = R"(
<mujoco>
<geom name="box" type="box" size="1 1 1"/>
</mujoco>)";
auto vfs = std::make_unique<mjVFS>();
mj_defaultVFS(vfs.get());
mj_makeEmptyFileVFS(vfs.get(), "model1.xml", sizeof(xml1));
std::memcpy(vfs->filedata[vfs->nfile - 1], xml1, sizeof(xml1));
std::array<char, 1024> 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"(
<mujoco>
@@ -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"(