diff --git a/src/xml/xml.cc b/src/xml/xml.cc index e87630bd..7b8ec3cd 100644 --- a/src/xml/xml.cc +++ b/src/xml/xml.cc @@ -92,37 +92,18 @@ class LocaleOverride { } // namespace // Main writer function - calls mjXWrite -bool mjWriteXML(mjCModel* model, string filename, char* error, int error_sz) { +string mjWriteXML(mjCModel* model, char* error, int error_sz) { LocaleOverride locale_override; // check for empty model if (!model) { mjCopyError(error, "Cannot write empty model", error_sz); - return false; + return ""; } - // write - FILE* fp = fopen(filename.c_str(), "w"); - if (!fp) { - mjCopyError(error, "File not found", error_sz); - return false; - } - - try { - mjXWriter writer; - writer.SetModel(model); - writer.Write(fp); - } - - // catch known errors - catch (mjXError err) { - mjCopyError(error, err.message, error_sz); - fclose(fp); - return false; - } - - fclose(fp); - return true; + mjXWriter writer; + writer.SetModel(model); + return writer.Write(error, error_sz); } diff --git a/src/xml/xml.h b/src/xml/xml.h index 8b39ed35..acbe81ce 100644 --- a/src/xml/xml.h +++ b/src/xml/xml.h @@ -23,7 +23,7 @@ // Top level API // Main writer function -bool mjWriteXML(mjCModel* model, std::string filename, char* error, int error_sz); +std::string mjWriteXML(mjCModel* model, char* error, int error_sz); // Main parser function mjCModel* mjParseXML(const char* filename, int vfs_provider, char* error, int error_sz); diff --git a/src/xml/xml_api.cc b/src/xml/xml_api.cc index 39dce2ee..bad07392 100644 --- a/src/xml/xml_api.cc +++ b/src/xml/xml_api.cc @@ -137,24 +137,37 @@ mjModel* mj_loadXML(const char* filename, const mjVFS* vfs, int mj_saveLastXML(const char* filename, const mjModel* m, char* error, int error_sz) { // serialize access to themodel std::lock_guard lock(themutex); + FILE *fp = stdout; if (!themodel.model) { mjCopyError(error, "No XML model loaded", error_sz); return 0; } - themodel.model->CopyBack(m); - if (mjWriteXML(themodel.model, filename, error, error_sz)) { - if (error) { - error[0] = 0; + if (filename != nullptr && filename[0] != '\0') { + fp = fopen(filename, "w"); + if (!fp) { + mjCopyError(error, "File not found", error_sz); + return 0; } - return 1; - } else { - return 0; } + + themodel.model->CopyBack(m); + std::string result = mjWriteXML(themodel.model, error, error_sz); + + if (!result.empty()) { + fprintf(fp, "%s", result.c_str()); + } + + if (fp != stdout) { + fclose(fp); + } + + return !result.empty(); } + // free last XML void mj_freeLastXML(void) { // serialize access to themodel diff --git a/src/xml/xml_base.h b/src/xml/xml_base.h index 4e971567..1b2c5c6b 100644 --- a/src/xml/xml_base.h +++ b/src/xml/xml_base.h @@ -15,6 +15,7 @@ #ifndef MUJOCO_SRC_XML_XML_BASE_H_ #define MUJOCO_SRC_XML_XML_BASE_H_ +#include #include #include "tinyxml2.h" @@ -82,7 +83,9 @@ class mjXBase : public mjXUtil { virtual void Parse(tinyxml2::XMLElement* root) {}; // write: implemented in derived writer class - virtual void Write(FILE* fp) {}; + virtual std::string Write(char *error, std::size_t error_sz) { + return ""; + }; // set the model allocated externally void SetModel(mjCModel*); diff --git a/src/xml/xml_native_writer.cc b/src/xml/xml_native_writer.cc index 91d02bb2..00abbe16 100644 --- a/src/xml/xml_native_writer.cc +++ b/src/xml/xml_native_writer.cc @@ -59,11 +59,15 @@ class mj_XMLPrinter : public tinyxml2::XMLPrinter { // save XML file using custom 2-space indentation -tinyxml2::XMLError SaveFile(XMLDocument& doc, FILE* fp) { +static string WriteDoc(XMLDocument& doc, char *error, size_t error_sz) { doc.ClearError(); - mj_XMLPrinter stream(fp, /*compact=*/false); + mj_XMLPrinter stream(nullptr, /*compact=*/false); doc.Print(&stream); - return doc.ErrorID(); + if (doc.ErrorID()) { + mjCopyError(error, doc.ErrorStr(), error_sz); + return ""; + } + return string(stream.CStr()); } @@ -682,10 +686,11 @@ mjXWriter::mjXWriter(void) { // save existing model in MJCF canonical format, must be compiled -void mjXWriter::Write(FILE* fp) { +string mjXWriter::Write(char *error, size_t error_sz) { // check model if (!model || !model->IsCompiled()) { - throw mjXError(0, "XML Write error: Only compiled model can be written"); + mjCopyError(error, "XML Write error: Only compiled model can be written", error_sz); + return ""; } // create document and root @@ -722,8 +727,7 @@ void mjXWriter::Write(FILE* fp) { Sensor(root); Keyframe(root); - // save file - SaveFile(doc, fp); + return WriteDoc(doc, error, error_sz); } diff --git a/src/xml/xml_native_writer.h b/src/xml/xml_native_writer.h index 862c4bb1..c5da1d87 100644 --- a/src/xml/xml_native_writer.h +++ b/src/xml/xml_native_writer.h @@ -15,6 +15,7 @@ #ifndef MUJOCO_SRC_XML_XML_NATIVE_WRITER_H_ #define MUJOCO_SRC_XML_XML_NATIVE_WRITER_H_ +#include #include #include "xml/xml_base.h" @@ -24,7 +25,9 @@ class mjXWriter : public mjXBase { public: mjXWriter(); // constructor virtual ~mjXWriter() = default; // destructor - void Write(FILE* fp); // write XML document + + // write XML document to string + std::string Write(char *error, std::size_t error_sz); private: // insert end child with given name, return child diff --git a/test/xml/xml_api_test.cc b/test/xml/xml_api_test.cc index 1c5e87b4..ade892f9 100644 --- a/test/xml/xml_api_test.cc +++ b/test/xml/xml_api_test.cc @@ -14,6 +14,7 @@ // Tests for xml/xml_api.cc. +#include #include #include #include @@ -30,6 +31,7 @@ namespace { using ::testing::IsNull; using ::testing::NotNull; +using ::testing::StartsWith; // ---------------------------- test mj_loadXML -------------------------------- @@ -54,16 +56,34 @@ TEST_F(LoadXmlTest, EmptyModel) { TEST_F(LoadXmlTest, InvalidXmlFailsToLoad) { static constexpr char invalid_xml[] = " error; + mjModel* model = LoadModelFromString(invalid_xml, error.data(), error.size()); EXPECT_THAT(model, IsNull()) << "Expected model loading to fail."; - EXPECT_GT(std::strlen(error), 0); + EXPECT_GT(std::strlen(error.data()), 0); if (model) { mj_deleteModel(model); } } // TODO(nimrod): Add more tests for mj_loadXML. +using SaveLastXmlTest = MujocoTest; + +TEST_F(SaveLastXmlTest, EmptyModel) { + static constexpr char xml[] = ""; + mjModel* model = LoadModelFromString(xml, 0, 0); + mjData* data = mj_makeData(model); + + std::array error; + error.data()[0] = '\0'; + + testing::internal::CaptureStdout(); + mj_saveLastXML(nullptr, model, error.data(), error.size()); + + EXPECT_THAT(testing::internal::GetCapturedStdout(), StartsWith("