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
+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