Add a new plugin / extension mechanism called a resource provider along with retrofitting VFS on top of it.

A resource provider provides a mechanism for MuJoCo to read from filesystems other than the OS filesystem or the Virtual File System (VFS).

PiperOrigin-RevId: 525394983
Change-Id: I077ff5a7e2e76806b48b6defb531280aadc8b169
This commit is contained in:
Kyle Bayes
2023-04-19 03:01:33 -07:00
committed by Copybara-Service
parent b25728cc2e
commit fe3dccfd1d
27 changed files with 1467 additions and 607 deletions
+61 -37
View File
@@ -24,7 +24,7 @@
#include "cc/array_safety.h"
#include "engine/engine_crossplatform.h"
#include "engine/engine_vfs.h"
#include "engine/engine_resource.h"
#include "user/user_model.h"
#include "user/user_util.h"
#include "xml/xml_native_reader.h"
@@ -129,7 +129,7 @@ bool mjWriteXML(mjCModel* model, string filename, 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) {
int default_provider, vector<string>& included) {
// include element: process
if (!strcasecmp(elem->Value(), "include")) {
// make sure include has no children
@@ -150,23 +150,30 @@ static XMLElement* mjIncludeXML(XMLElement* elem, string dir,
}
// get data source
const char* xmlstring = 0;
int buffer_size = 0;
if (vfs) {
int id = mj_findFileVFS(vfs, filename.c_str());
if (id>=0) {
xmlstring = (const char*)vfs->filedata[id];
buffer_size = vfs->filesize[id];
mjResource *resource = nullptr;
const char* xmlstring = nullptr;
if ((resource = mju_openResource(filename.c_str(), default_provider)) == nullptr) {
// load from OS filesystem
if (!default_provider || (resource = mju_openResource(filename.c_str(), 0)) == 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;
if (xmlstring) {
doc.Parse(xmlstring, buffer_size);
} else {
doc.LoadFile(filename.c_str());
}
doc.Parse(xmlstring, buffer_size);
// close resource
mju_closeResource(resource);
// check error
if (doc.Error()) {
@@ -208,27 +215,26 @@ static XMLElement* mjIncludeXML(XMLElement* elem, string dir,
}
// run XMLInclude on first new child
return mjIncludeXML(first->ToElement(), dir, vfs, included);
return mjIncludeXML(first->ToElement(), dir, default_provider, included);
}
// otherwise check all child elements, return self
else {
XMLElement* child = elem->FirstChildElement();
while (child) {
child = mjIncludeXML(child, dir, vfs, included);
child = mjIncludeXML(child, dir, default_provider, included);
if (child) {
child = child->NextSiblingElement();
}
}
return elem;
}
}
// Main parser function: from file or VFS
mjCModel* mjParseXML(const char* filename, const mjVFS* vfs, char* error, int error_sz) {
// Main parser function
mjCModel* mjParseXML(const char* filename, int default_provider, char* error, int error_sz) {
LocaleOverride locale_override;
// check arguments
@@ -236,33 +242,51 @@ mjCModel* mjParseXML(const char* filename, const mjVFS* vfs, char* error, int er
if (error) {
snprintf(error, error_sz, "mjParseXML: filename argument required\n");
}
return 0;
return nullptr;
}
// clear
mjCModel* model = 0;
if (error) {
error[0] = 0;
error[0] = '\0';
}
// get data source
const char* xmlstring = 0;
int buffer_size = 0;
if (vfs) {
int id = mj_findFileVFS(vfs, filename);
if (id>=0) {
xmlstring = (const char*)vfs->filedata[id];
buffer_size = vfs->filesize[id];
mjResource* resource = nullptr;
const char* xmlstring = nullptr;
if ((resource = mju_openResource(filename, default_provider)) == nullptr) {
// load from OS filesystem
if (!default_provider || (resource = mju_openResource(filename, 0)) == nullptr) {
if (error) {
snprintf(error, error_sz, "mjParseXML: could not open file '%s'", filename);
}
return nullptr;
}
}
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);
}
mju_closeResource(resource);
return nullptr;
} else if (!buffer_size) {
if (error) {
snprintf(error, error_sz, "mjParseXML: empty file '%s'", filename);
}
mju_closeResource(resource);
return nullptr;
}
// load XML file or parse string
XMLDocument doc;
if (xmlstring) {
doc.Parse(xmlstring, buffer_size);
} else {
doc.LoadFile(filename);
}
doc.Parse(xmlstring, buffer_size);
// close resource
mju_closeResource(resource);
// error checking
if (doc.Error()) {
@@ -270,14 +294,14 @@ mjCModel* mjParseXML(const char* filename, const mjVFS* vfs, char* error, int er
snprintf(error, error_sz, "XML parse error %d:\n%s\n",
doc.ErrorID(), doc.ErrorStr());
}
return 0;
return nullptr;
}
// get top-level element
XMLElement* root = doc.RootElement();
if (!root) {
mjCopyError(error, "XML root element not found", error_sz);
return 0;
return nullptr;
}
// create model, set filedir
@@ -290,7 +314,7 @@ mjCModel* mjParseXML(const char* filename, const mjVFS* vfs, char* error, int er
// find include elements, replace them with subtree from xml file
vector<string> included;
included.push_back(filename);
mjIncludeXML(root, model->modelfiledir, vfs, included);
mjIncludeXML(root, model->modelfiledir, default_provider, included);
// parse MuJoCo model
mjXReader parser;
@@ -314,7 +338,7 @@ mjCModel* mjParseXML(const char* filename, const mjVFS* vfs, char* error, int er
catch (mjXError err) {
mjCopyError(error, err.message, error_sz);
delete model;
return 0;
return nullptr;
}
return model;
+2 -2
View File
@@ -25,8 +25,8 @@
// Main writer function
bool mjWriteXML(mjCModel* model, std::string filename, char* error, int error_sz);
// Main parser function: from file or VFS
mjCModel* mjParseXML(const char* filename, const mjVFS* vfs, char* error, int error_sz);
// Main parser function
mjCModel* mjParseXML(const char* filename, int default_provider, char* error, int error_sz);
#endif // MUJOCO_SRC_XML_XML_H_
+34 -10
View File
@@ -20,6 +20,8 @@
#include <mutex>
#include <random>
#include "engine/engine_resource.h"
#include "engine/engine_vfs.h"
#include "user/user_model.h"
#include "xml/xml.h"
#include "xml/xml_native_reader.h"
@@ -80,26 +82,24 @@ void mj_deactivate(void) {
// parse XML file in MJCF or URDF format, compile it, return low-level model
// if vfs is not NULL, look up files in vfs before reading from disk
// error can be NULL; otherwise assumed to have size error_sz
mjModel* mj_loadXML(const char* filename, const mjVFS* vfs,
char* error, int error_sz) {
// mj_loadXML helper function
mjModel* _loadXML(const char* filename, int default_provider,
char* error, int error_sz) {
// serialize access to themodel
std::lock_guard<std::mutex> lock(themutex);
// parse new model
mjCModel* newmodel = mjParseXML(filename, vfs, error, error_sz);
mjCModel* newmodel = mjParseXML(filename, default_provider, error, error_sz);
if (!newmodel) {
return 0;
return nullptr;
}
// compile new model
mjModel* m = newmodel->Compile(vfs);
mjModel* m = newmodel->Compile(default_provider);
if (!m) {
mjCopyError(error, newmodel->GetError().message, error_sz);
delete newmodel;
return 0;
return nullptr;
}
// clear old and assign new
@@ -110,7 +110,7 @@ mjModel* mj_loadXML(const char* filename, const mjVFS* vfs,
if (themodel.model->GetError().warning) {
mjCopyError(error, themodel.model->GetError().message, error_sz);
} else if (error) {
error[0] = 0;
error[0] = '\0';
}
return m;
@@ -118,6 +118,30 @@ mjModel* mj_loadXML(const char* filename, const mjVFS* vfs,
// parse XML file in MJCF or URDF format, compile it, return low-level model
// if vfs is not NULL, look up files in vfs before reading from disk
// error can be NULL; otherwise assumed to have size error_sz
mjModel* mj_loadXML(const char* filename, const mjVFS* vfs,
char* error, int error_sz) {
if (vfs == nullptr) {
return _loadXML(filename, 0, error, error_sz);
}
int index = mj_registerVfsProvider(vfs);
if (index < 1) {
if (error) {
snprintf(error, error_sz, "mj_loadXML: could not register VFS");
}
return nullptr;
}
mjModel* model = _loadXML(filename, index, error, error_sz);
mjp_unregisterResourceProvider(index);
return model;
}
// update XML data structures with info from low-level model, save as MJCF
// returns 1 if successful, 0 otherwise