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