From b35ae973e19e1793963c5d15571c7adafa640d63 Mon Sep 17 00:00:00 2001 From: Alessio Quaglino Date: Tue, 4 Jun 2024 11:59:32 +0100 Subject: [PATCH] Add mj_parseXML and mj_saveXML to xml_api. PiperOrigin-RevId: 640099466 Change-Id: I75f763523c5bc8f6f7c2f2de6806d82ba37ebd4e --- src/xml/xml.cc | 6 +-- src/xml/xml.h | 7 ++- src/xml/xml_api.cc | 50 +++++++++++++++++++ src/xml/xml_api.h | 9 ++++ src/xml/xml_base.cc | 4 +- src/xml/xml_base.h | 2 +- src/xml/xml_native_writer.cc | 2 +- src/xml/xml_native_writer.h | 2 +- test/user/user_api_test.cc | 16 +++--- test/xml/xml_api_test.cc | 78 ++++++++++++++++++++++++------ test/xml/xml_native_reader_test.cc | 4 +- 11 files changed, 143 insertions(+), 37 deletions(-) diff --git a/src/xml/xml.cc b/src/xml/xml.cc index dfd60c61..ad1cf06f 100644 --- a/src/xml/xml.cc +++ b/src/xml/xml.cc @@ -100,7 +100,7 @@ class LocaleOverride { } // namespace // Main writer function - calls mjXWrite -std::string mjWriteXML(mjSpec* spec, char* error, int error_sz) { +std::string mjWriteXML(const mjSpec* spec, char* error, int error_sz) { LocaleOverride locale_override; // check for empty model @@ -424,9 +424,9 @@ static void RegisterResourceProvider() { mjSpec* ParseSpecFromString(std::string_view xml, char* error, - int error_size, mjVFS* vfs) { + int error_size) { RegisterResourceProvider(); std::string xml2 = {xml.begin(), xml.end()}; std::string str = "LoadModelFromString:" + xml2; - return mjParseXML(str.c_str(), vfs, error, error_size); + return mjParseXML(str.c_str(), nullptr, error, error_size); } diff --git a/src/xml/xml.h b/src/xml/xml.h index 10b5e4a3..2fdd3acc 100644 --- a/src/xml/xml.h +++ b/src/xml/xml.h @@ -24,15 +24,14 @@ // Top level API // Main writer function -std::string mjWriteXML(mjSpec* spec, char* error, int error_sz); +std::string mjWriteXML(const mjSpec* spec, char* error, int error_sz); // Main parser function -MJAPI mjSpec* mjParseXML(const char* filename, const mjVFS* vfs, char* error, int error_sz); +mjSpec* mjParseXML(const char* filename, const mjVFS* vfs, char* error, int error_sz); // Returns a newly-allocated mjSpec, loaded from the contents of xml. // On failure returns nullptr and populates the error array if present. -MJAPI mjSpec* ParseSpecFromString(std::string_view xml, char* error = nullptr, - int error_size = 0, mjVFS* vfs = nullptr); +mjSpec* ParseSpecFromString(std::string_view xml, char* error = nullptr, int error_size = 0); #endif // MUJOCO_SRC_XML_XML_H_ diff --git a/src/xml/xml_api.cc b/src/xml/xml_api.cc index 1f30e8a5..23af4c81 100644 --- a/src/xml/xml_api.cc +++ b/src/xml/xml_api.cc @@ -209,3 +209,53 @@ mjModel* mj_loadModel(const char* filename, const mjVFS* vfs) { return m; } + + +// parse spec from file +mjSpec* mj_parseXML(const char* filename, const mjVFS* vfs, char* error, int error_sz) { + return mjParseXML(filename, vfs, error, error_sz); +} + + + +// parse spec from string +mjSpec* mj_parseXMLString(const char* xml, const mjVFS* vfs, char* error, int error_sz) { + return ParseSpecFromString(xml, error, error_sz); +} + + + +// save spec to XML file, return 1 on success, 0 otherwise +int mj_saveXML(const mjSpec* s, const char* filename, char* error, int error_sz) { + std::string result = mjWriteXML(s, error, error_sz); + if (result.empty()) { + return 0; + } + + std::ofstream file; + file.open(filename); + file << result; + file.close(); + return 1; +} + + + +// save spec to string, return 1 on success, 0 otherwise +int mj_saveXMLString(const mjSpec* s, char* xml, int xml_sz, char* error, int error_sz) { + std::string result = mjWriteXML(s, error, error_sz); + if (result.size() >= xml_sz) { + std::string error_msg = "Output string too short, should be at least " + + std::to_string(result.size()+1); + mjCopyError(error, error_msg.c_str(), error_sz); + return 0; + } + if (result.empty()) { + return 0; + } + + result.copy(xml, xml_sz); + xml[result.size()] = 0; + return 1; +} + diff --git a/src/xml/xml_api.h b/src/xml/xml_api.h index bc2a5a71..046e8410 100644 --- a/src/xml/xml_api.h +++ b/src/xml/xml_api.h @@ -17,6 +17,7 @@ #include #include +#include "user/user_api.h" #ifdef __cplusplus extern "C" { @@ -43,6 +44,14 @@ MJAPI int mj_printSchema(const char* filename, char* buffer, int buffer_sz, // if vfs is not NULL, look up file in vfs before reading from disk MJAPI mjModel* mj_loadModel(const char* filename, const mjVFS* vfs); +// parse spec from file or XML string. +MJAPI mjSpec* mj_parseXML(const char* filename, const mjVFS* vfs, char* error, int error_sz); +MJAPI mjSpec* mj_parseXMLString(const char* xml, const mjVFS* vfs, char* error, int error_sz); + +// Save spec to XML file and/or string, return 1 on success, 0 otherwise. +MJAPI int mj_saveXML(const mjSpec* s, const char* filename, char* error, int error_sz); +MJAPI int mj_saveXMLString(const mjSpec* s, char* xml, int xml_sz, char* error, int error_sz); + #ifdef __cplusplus } #endif diff --git a/src/xml/xml_base.cc b/src/xml/xml_base.cc index 307dc69e..bcefd583 100644 --- a/src/xml/xml_base.cc +++ b/src/xml/xml_base.cc @@ -45,8 +45,8 @@ mjXBase::mjXBase() { // set model field -void mjXBase::SetModel(mjSpec* _model) { - model = _model; +void mjXBase::SetModel(const mjSpec* _model) { + model = (mjSpec*)_model; } diff --git a/src/xml/xml_base.h b/src/xml/xml_base.h index 26d95fcb..7df966cc 100644 --- a/src/xml/xml_base.h +++ b/src/xml/xml_base.h @@ -87,7 +87,7 @@ class mjXBase : public mjXUtil { }; // set the model allocated externally - virtual void SetModel(mjSpec*); + virtual void SetModel(const mjSpec*); // read alternative orientation specification static int ReadAlternative(tinyxml2::XMLElement* elem, mjsOrientation& alt); diff --git a/src/xml/xml_native_writer.cc b/src/xml/xml_native_writer.cc index 97daccd4..c9012bdd 100644 --- a/src/xml/xml_native_writer.cc +++ b/src/xml/xml_native_writer.cc @@ -769,7 +769,7 @@ mjXWriter::mjXWriter(void) { // cast model -void mjXWriter::SetModel(mjSpec* spec) { +void mjXWriter::SetModel(const mjSpec* spec) { if (spec) { model = (mjCModel*)spec->element; } diff --git a/src/xml/xml_native_writer.h b/src/xml/xml_native_writer.h index 063fddca..6f5f0812 100644 --- a/src/xml/xml_native_writer.h +++ b/src/xml/xml_native_writer.h @@ -27,7 +27,7 @@ class mjXWriter : public mjXBase { public: mjXWriter(); // constructor virtual ~mjXWriter() = default; // destructor - void SetModel(mjSpec* spec); + void SetModel(const mjSpec* spec); // write XML document to string std::string Write(char *error, std::size_t error_sz); diff --git a/test/user/user_api_test.cc b/test/user/user_api_test.cc index 440f0c96..e6b63761 100644 --- a/test/user/user_api_test.cc +++ b/test/user/user_api_test.cc @@ -25,7 +25,7 @@ #include #include #include "src/user/user_api.h" -#include "src/xml/xml.h" +#include "src/xml/xml_api.h" #include "src/xml/xml_numeric_format.h" #include "test/fixture.h" @@ -125,7 +125,7 @@ TEST_F(PluginTest, RecompileCompare) { // load spec std::array err; - mjSpec* s = mjParseXML(xml.c_str(), nullptr, err.data(), err.size()); + mjSpec* s = mj_parseXML(xml.c_str(), 0, err.data(), err.size()); ASSERT_THAT(s, NotNull()) << "Failed to load " << xml << ": " << err.data(); @@ -417,7 +417,7 @@ TEST_F(MujocoTest, AttachSame) { )"; // create parent - mjSpec* parent = ParseSpecFromString(xml_child, er.data(), er.size()); + mjSpec* parent = mj_parseXMLString(xml_child, 0, er.data(), er.size()); EXPECT_THAT(parent, NotNull()) << er.data(); // get frame @@ -530,7 +530,7 @@ TEST_F(MujocoTest, AttachDifferent) { )"; // model with one free sphere and a frame - mjSpec* parent = ParseSpecFromString(xml_parent, er.data(), er.size()); + mjSpec* parent = mj_parseXMLString(xml_parent, 0, er.data(), er.size()); EXPECT_THAT(parent, NotNull()) << er.data(); // get frame @@ -538,7 +538,7 @@ TEST_F(MujocoTest, AttachDifferent) { EXPECT_THAT(frame, NotNull()); // model with one cylinder and a hinge - mjSpec* child = ParseSpecFromString(xml_child, er.data(), er.size()); + mjSpec* child = mj_parseXMLString(xml_child, 0, er.data(), er.size()); EXPECT_THAT(child, NotNull()) << er.data(); // get subtree @@ -642,7 +642,7 @@ TEST_F(MujocoTest, AttachFrame) { )"; // model with one free sphere and a frame - mjSpec* parent = ParseSpecFromString(xml_parent, er.data(), er.size()); + mjSpec* parent = mj_parseXMLString(xml_parent, 0, er.data(), er.size()); EXPECT_THAT(parent, NotNull()) << er.data(); // get frame @@ -650,7 +650,7 @@ TEST_F(MujocoTest, AttachFrame) { EXPECT_THAT(body, NotNull()); // model with one cylinder and a hinge - mjSpec* child = ParseSpecFromString(xml_child, er.data(), er.size()); + mjSpec* child = mj_parseXMLString(xml_child, 0, er.data(), er.size()); EXPECT_THAT(child, NotNull()) << er.data(); // get subtree @@ -711,7 +711,7 @@ void TestDetachBody(bool compile) { )"; // model with one cylinder and a hinge - mjSpec* child = ParseSpecFromString(xml_child, er.data(), er.size()); + mjSpec* child = mj_parseXMLString(xml_child, 0, er.data(), er.size()); EXPECT_THAT(child, NotNull()) << er.data(); // compile model (for testing double compilation) diff --git a/test/xml/xml_api_test.cc b/test/xml/xml_api_test.cc index 8988b851..d11be846 100644 --- a/test/xml/xml_api_test.cc +++ b/test/xml/xml_api_test.cc @@ -24,6 +24,8 @@ #include #include #include +#include "src/user/user_api.h" +#include "src/xml/xml_api.h" #include "test/fixture.h" namespace mujoco { @@ -33,6 +35,21 @@ using ::testing::IsNull; using ::testing::NotNull; using ::testing::StartsWith; +static constexpr char xml[] = R"( + + + + + + + + + + + + + )"; + // ---------------------------- test mj_loadXML -------------------------------- using LoadXmlTest = MujocoTest; @@ -66,20 +83,6 @@ TEST_F(LoadXmlTest, InvalidXmlFailsToLoad) { } TEST_F(LoadXmlTest, MultipleBodies) { - static constexpr char xml[] = R"( - - - - - - - - - - - - - )"; std::array error; mjModel* model = LoadModelFromString(xml, error.data(), error.size()); @@ -92,7 +95,6 @@ TEST_F(LoadXmlTest, MultipleBodies) { mj_deleteData(data); mj_deleteModel(model); } - using SaveLastXmlTest = MujocoTest; TEST_F(SaveLastXmlTest, EmptyModel) { @@ -112,5 +114,51 @@ TEST_F(SaveLastXmlTest, EmptyModel) { mj_deleteModel(model); } +TEST_F(MujocoTest, SaveXmlShortString) { + std::array error; + + mjSpec* spec = mj_parseXMLString(xml, 0, error.data(), error.size()); + EXPECT_THAT(spec, NotNull()) << "Failed to parse spec: " << error.data(); + mjModel* model = mjs_compile(spec, 0); + EXPECT_THAT(model, NotNull()) << "Failed to compile model: " << error.data(); + + std::array out; + EXPECT_THAT(mj_saveXMLString(spec, out.data(), out.size(), + error.data(), error.size()), 0); + EXPECT_STREQ(error.data(), "Output string too short, should be at least 273"); + + mjs_deleteSpec(spec); + mj_deleteModel(model); +} + +TEST_F(MujocoTest, SaveXml) { + std::array error; + + mjSpec* spec = mj_parseXMLString(xml, 0, error.data(), error.size()); + EXPECT_THAT(spec, NotNull()) << "Failed to parse spec: " << error.data(); + mjModel* model = mjs_compile(spec, 0); + EXPECT_THAT(model, NotNull()) << "Failed to compile model: " << error.data(); + + std::array out; + EXPECT_THAT(mj_saveXMLString(spec, out.data(), out.size(), error.data(), + error.size()), 1) << error.data(); + + mjSpec* saved_spec = mj_parseXMLString(xml, 0, error.data(), error.size()); + EXPECT_THAT(saved_spec, NotNull()) << "Invalid saved spec: " << error.data(); + mjModel* saved_model = mjs_compile(saved_spec, 0); + EXPECT_THAT(saved_model, NotNull()) << "Invalid model: " << error.data(); + + mjtNum tol = 0; + std::string field = ""; + EXPECT_LE(CompareModel(model, saved_model, field), tol) + << "Expected and attached models are different!\n" + << "Different field: " << field << '\n'; + + mjs_deleteSpec(spec); + mjs_deleteSpec(saved_spec); + mj_deleteModel(model); + mj_deleteModel(saved_model); +} + } // namespace } // namespace mujoco diff --git a/test/xml/xml_native_reader_test.cc b/test/xml/xml_native_reader_test.cc index 1f225fa4..c07c8e73 100644 --- a/test/xml/xml_native_reader_test.cc +++ b/test/xml/xml_native_reader_test.cc @@ -28,7 +28,7 @@ #include "src/cc/array_safety.h" #include "src/engine/engine_util_errmem.h" #include "src/user/user_api.h" -#include "src/xml/xml.h" +#include "src/xml/xml_api.h" #include "test/fixture.h" namespace mujoco { @@ -1084,7 +1084,7 @@ TEST_F(XMLReaderTest, ParseReplicateDefaultPropagate) { )"; std::array error; - mjSpec* spec = ParseSpecFromString(xml, error.data(), error.size()); + mjSpec* spec = mj_parseXMLString(xml, 0, error.data(), error.size()); EXPECT_THAT(spec, NotNull()) << error.data(); mjsBody* torso = mjs_findBody(spec, "torso-0");