Add pluggable resource writing to MuJoCo
Extend mjpResourceProvider with an optional write callback (write) so that mj_encode, mj_saveXML, and mj_saveModel can write to any registered provider. PiperOrigin-RevId: 945202741 Change-Id: I37903425260932e555f4a8c2392c4ff8c2e6cc06
This commit is contained in:
committed by
Copybara-Service
parent
bdeb7e7c4d
commit
dc7581acfa
+25
-38
@@ -503,56 +503,43 @@ void mjv_copyModel(mjModel* dest, const mjModel* src) {
|
||||
|
||||
// save model to binary file, or memory buffer of szbuf>0
|
||||
void mj_saveModel(const mjModel* m, const char* filename, void* buffer, int buffer_sz) {
|
||||
FILE* fp = 0;
|
||||
mjtSize ptrbuf = 0;
|
||||
|
||||
// standard header
|
||||
int header[NHEADER] = {ID, sizeof(mjtNum), getnsize(), mj_version(), getnptr()};
|
||||
|
||||
// open file for writing if no buffer
|
||||
// no buffer: serialize to temporary buffer, then write via resource provider
|
||||
if (!buffer) {
|
||||
fp = fopen(filename, "wb");
|
||||
if (!fp) {
|
||||
mju_warning("Could not open file '%s'", filename);
|
||||
mjtSize sz = mj_sizeModel(m);
|
||||
void* tmpbuf = mju_malloc(sz);
|
||||
if (!tmpbuf) {
|
||||
mju_warning("Could not allocate buffer for saving model");
|
||||
return;
|
||||
}
|
||||
mj_saveModel(m, NULL, tmpbuf, (int)sz);
|
||||
|
||||
mjtSize written = mju_writeResource(filename, tmpbuf, sz, NULL, NULL, 0);
|
||||
if (written != sz) {
|
||||
mju_warning("Could not save model to '%s'", filename);
|
||||
}
|
||||
mju_free(tmpbuf);
|
||||
return;
|
||||
}
|
||||
|
||||
// write standard header, info, options, buffer (omit pointers)
|
||||
if (fp) {
|
||||
fwrite(header, sizeof(int), NHEADER, fp);
|
||||
#define X(name) fwrite(&m->name, sizeof(m->name), 1, fp);
|
||||
MJMODEL_SIZES
|
||||
bufwrite(header, sizeof(header), buffer_sz, buffer, &ptrbuf);
|
||||
#define X(name) bufwrite(&m->name, sizeof(m->name), buffer_sz, buffer, &ptrbuf);
|
||||
MJMODEL_SIZES
|
||||
#undef X
|
||||
bufwrite((void*)&m->opt, sizeof(mjOption), buffer_sz, buffer, &ptrbuf);
|
||||
bufwrite((void*)&m->vis, sizeof(mjVisual), buffer_sz, buffer, &ptrbuf);
|
||||
bufwrite((void*)&m->stat, sizeof(mjStatistic), buffer_sz, buffer, &ptrbuf);
|
||||
{
|
||||
MJMODEL_POINTERS_PREAMBLE(m)
|
||||
#define X(type, name, nr, nc) \
|
||||
bufwrite((void*)m->name, sizeof(type)*(m->nr)*(nc), buffer_sz, buffer, &ptrbuf);
|
||||
MJMODEL_POINTERS
|
||||
#undef X
|
||||
fwrite((void*)&m->opt, sizeof(mjOption), 1, fp);
|
||||
fwrite((void*)&m->vis, sizeof(mjVisual), 1, fp);
|
||||
fwrite((void*)&m->stat, sizeof(mjStatistic), 1, fp);
|
||||
{
|
||||
MJMODEL_POINTERS_PREAMBLE(m)
|
||||
#define X(type, name, nr, nc) \
|
||||
fwrite((void*)m->name, sizeof(type), (m->nr)*(nc), fp);
|
||||
MJMODEL_POINTERS
|
||||
#undef X
|
||||
}
|
||||
} else {
|
||||
bufwrite(header, sizeof(header), buffer_sz, buffer, &ptrbuf);
|
||||
#define X(name) bufwrite(&m->name, sizeof(m->name), buffer_sz, buffer, &ptrbuf);
|
||||
MJMODEL_SIZES
|
||||
#undef X
|
||||
bufwrite((void*)&m->opt, sizeof(mjOption), buffer_sz, buffer, &ptrbuf);
|
||||
bufwrite((void*)&m->vis, sizeof(mjVisual), buffer_sz, buffer, &ptrbuf);
|
||||
bufwrite((void*)&m->stat, sizeof(mjStatistic), buffer_sz, buffer, &ptrbuf);
|
||||
{
|
||||
MJMODEL_POINTERS_PREAMBLE(m)
|
||||
#define X(type, name, nr, nc) \
|
||||
bufwrite((void*)m->name, sizeof(type)*(m->nr)*(nc), buffer_sz, buffer, &ptrbuf);
|
||||
MJMODEL_POINTERS
|
||||
#undef X
|
||||
}
|
||||
}
|
||||
|
||||
if (fp) {
|
||||
fclose(fp);
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -249,11 +249,10 @@ std::string_view GlobalTable<mjpResourceProvider>::ObjectKey(const mjpResourcePr
|
||||
// check if two resource providers are identical
|
||||
template <>
|
||||
bool GlobalTable<mjpResourceProvider>::ObjectEqual(const mjpResourceProvider& p1, const mjpResourceProvider& p2) {
|
||||
return (CaseInsensitiveEqual(p1.prefix, p2.prefix) &&
|
||||
p1.open == p2.open &&
|
||||
p1.read == p2.read &&
|
||||
p1.close == p2.close &&
|
||||
return (CaseInsensitiveEqual(p1.prefix, p2.prefix) && p1.open == p2.open &&
|
||||
p1.read == p2.read && p1.close == p2.close &&
|
||||
p1.modified == p2.modified &&
|
||||
p1.write == p2.write &&
|
||||
p1.data == p2.data);
|
||||
}
|
||||
|
||||
@@ -390,7 +389,8 @@ bool GlobalTable<mjpEncoder>::ObjectEqual(const mjpEncoder& e1,
|
||||
}
|
||||
|
||||
return content_type_match && extension_match
|
||||
&& e1.encode == e2.encode && e1.close_resource == e2.close_resource;
|
||||
&& e1.encode == e2.encode
|
||||
&& e1.close_resource == e2.close_resource;
|
||||
}
|
||||
|
||||
template <>
|
||||
|
||||
+5
-16
@@ -261,29 +261,18 @@ mjtSize mj_encode(const mjSpec* s, const mjModel* m, const char* filename,
|
||||
return -1;
|
||||
}
|
||||
|
||||
FILE* fp = fopen(filename, "wb");
|
||||
if (!fp) {
|
||||
std::free(resource.data);
|
||||
if (error) {
|
||||
strncpy(error, "could not open file for writing", error_sz);
|
||||
error[error_sz - 1] = '\0';
|
||||
}
|
||||
return -1;
|
||||
}
|
||||
|
||||
const std::size_t written = fwrite(resource.data, 1, nbytes, fp);
|
||||
fclose(fp);
|
||||
mjtSize written = mju_writeResource(filename, resource.data, nbytes, vfs, error, error_sz);
|
||||
encoder->close_resource(&resource);
|
||||
|
||||
if (static_cast<mjtSize>(written) != nbytes) {
|
||||
if (error) {
|
||||
strncpy(error, "failed to write all bytes to file", error_sz);
|
||||
if (written != nbytes) {
|
||||
if (error && error[0] == '\0') {
|
||||
strncpy(error, "write failed", error_sz);
|
||||
error[error_sz - 1] = '\0';
|
||||
}
|
||||
return -1;
|
||||
}
|
||||
|
||||
return nbytes;
|
||||
return written;
|
||||
}
|
||||
|
||||
// helper function to log compile time diagnostics
|
||||
|
||||
+56
-10
@@ -30,8 +30,13 @@
|
||||
#include "user/user_util.h"
|
||||
#include "user/user_vfs.h"
|
||||
|
||||
mjResource* mju_openResource(const char* dir, const char* name,
|
||||
const mjVFS* vfs, char* error, size_t nerror) {
|
||||
static mjResource* openResourceInternal(const char* dir, const char* name,
|
||||
const mjVFS* vfs, char* error,
|
||||
size_t nerror) {
|
||||
if (error && nerror > 0) {
|
||||
error[0] = '\0';
|
||||
}
|
||||
|
||||
// TODO: Update API to use non-const pointer. Unfortunately, while this is
|
||||
// ABI stable, it will cause compiler errors in user code that is const
|
||||
// correct.
|
||||
@@ -52,20 +57,21 @@ mjResource* mju_openResource(const char* dir, const char* name,
|
||||
non_const_vfs = local_vfs;
|
||||
}
|
||||
|
||||
mjResource* resource =
|
||||
mujoco::user::VFS::Upcast(non_const_vfs)->Open(dir ? dir : "", name);
|
||||
mjResource* resource = mujoco::user::VFS::Upcast(non_const_vfs)
|
||||
->Open(dir ? dir : "", name, error, nerror);
|
||||
|
||||
if (error) {
|
||||
if (resource) {
|
||||
error[0] = '\0';
|
||||
} else {
|
||||
std::snprintf(error, nerror, "Error opening file '%s'", name);
|
||||
}
|
||||
if (error && nerror > 0 && !resource && error[0] == '\0') {
|
||||
std::snprintf(error, nerror, "Error opening file '%s'", name);
|
||||
}
|
||||
|
||||
return resource;
|
||||
}
|
||||
|
||||
mjResource* mju_openResource(const char* dir, const char* name,
|
||||
const mjVFS* vfs, char* error, size_t nerror) {
|
||||
return openResourceInternal(dir, name, vfs, error, nerror);
|
||||
}
|
||||
|
||||
void mju_closeResource(mjResource* resource) {
|
||||
if (resource && resource->vfs) {
|
||||
mujoco::user::VFS::Upcast(resource->vfs)->Close(resource);
|
||||
@@ -79,6 +85,46 @@ int mju_readResource(mjResource* resource, const void** buffer) {
|
||||
return -1; // default (error reading bytes)
|
||||
}
|
||||
|
||||
mjtSize mju_writeResource(const char* name, const void* buffer, mjtSize nbytes,
|
||||
const mjVFS* vfs, char* error, size_t nerror) {
|
||||
if (error && nerror > 0) {
|
||||
error[0] = '\0';
|
||||
}
|
||||
|
||||
if (!name) {
|
||||
if (error && nerror > 0) {
|
||||
std::snprintf(error, nerror, "Resource name is NULL");
|
||||
}
|
||||
return -1;
|
||||
}
|
||||
|
||||
mjResource resource = {};
|
||||
resource.name = const_cast<char*>(name);
|
||||
|
||||
mjVFS* non_const_vfs = const_cast<mjVFS*>(vfs);
|
||||
bool local_vfs_created = false;
|
||||
if (non_const_vfs == nullptr) {
|
||||
non_const_vfs = (mjVFS*)mju_malloc(sizeof(mjVFS));
|
||||
mj_defaultVFS(non_const_vfs);
|
||||
local_vfs_created = true;
|
||||
}
|
||||
|
||||
resource.vfs = non_const_vfs;
|
||||
|
||||
mjtSize written = mujoco::user::VFS::Upcast(non_const_vfs)->Write(&resource, buffer, nbytes);
|
||||
|
||||
if (local_vfs_created) {
|
||||
mj_deleteVFS(non_const_vfs);
|
||||
mju_free(non_const_vfs);
|
||||
}
|
||||
|
||||
if (written < 0 && error && nerror > 0 && error[0] == '\0') {
|
||||
std::snprintf(error, nerror, "Error writing resource '%s'", name);
|
||||
}
|
||||
|
||||
return written;
|
||||
}
|
||||
|
||||
void mju_getResourceDir(mjResource* resource, const char** dir, int* ndir) {
|
||||
*dir = nullptr;
|
||||
*ndir = 0;
|
||||
|
||||
@@ -38,6 +38,9 @@ MJAPI void mju_closeResource(mjResource* resource);
|
||||
// return negative value if error
|
||||
MJAPI int mju_readResource(mjResource* resource, const void** buffer);
|
||||
|
||||
MJAPI mjtSize mju_writeResource(const char* name, const void* buffer, mjtSize nbytes,
|
||||
const mjVFS* vfs, char* error, size_t nerror);
|
||||
|
||||
// set for a resource with a name partitioned as {dir}{filename}, the dir and ndir pointers
|
||||
MJAPI void mju_getResourceDir(mjResource* resource, const char** dir, int* ndir);
|
||||
|
||||
|
||||
+39
-4
@@ -23,6 +23,7 @@
|
||||
#include <cctype>
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstdio>
|
||||
#include <cstring>
|
||||
#include <ctime>
|
||||
#include <functional>
|
||||
@@ -120,6 +121,15 @@ VFS::VFS(mjVFS* vfs) {
|
||||
default_provider_.modified = [](const mjResource* res, const char* time) {
|
||||
return FileModified(res, time);
|
||||
};
|
||||
default_provider_.write = [](mjResource* res, const void* buffer, mjtSize nbytes) -> mjtSize {
|
||||
FILE* fp = fopen(res->name, "wb");
|
||||
if (!fp) {
|
||||
return -1;
|
||||
}
|
||||
mjtSize written = static_cast<mjtSize>(fwrite(buffer, 1, nbytes, fp));
|
||||
fclose(fp);
|
||||
return written == nbytes ? written : -1;
|
||||
};
|
||||
|
||||
default_provider_.prefix = nullptr;
|
||||
default_mount_.vfs = &wrapped_vfs_;
|
||||
@@ -149,22 +159,28 @@ VFS::~VFS() {
|
||||
mounts_.clear();
|
||||
}
|
||||
|
||||
mjResource* VFS::Open(const char* dir, const char* name) {
|
||||
mjResource* VFS::Open(const char* dir, const char* name, char* error,
|
||||
size_t nerror) {
|
||||
const std::string path = FilePath(dir, name).Str();
|
||||
|
||||
const mjResource* mount = FindMount(path);
|
||||
if (!mount) {
|
||||
if (!mount || !mount->provider) {
|
||||
if (error) {
|
||||
std::snprintf(error, nerror, "No provider found for '%s'", name);
|
||||
}
|
||||
MaybeSelfDestruct();
|
||||
return nullptr;
|
||||
}
|
||||
|
||||
ResourcePtr res = CreateResource(path.c_str(), mount->provider);
|
||||
const mjpResourceProvider* provider = mount->provider;
|
||||
|
||||
ResourcePtr res = CreateResource(path.c_str(), provider);
|
||||
|
||||
// Smuggle the mounted resource provider's mount-specific data pointer in the
|
||||
// requested resource's data pointer. This allows the provider to access its
|
||||
// own per-mount data without any intrusive changes to the provider interface.
|
||||
res->data = mount->data;
|
||||
const int result = mount->provider->open(res.get());
|
||||
const int result = provider->open ? provider->open(res.get()) : 0;
|
||||
// If the data pointer was not modified, then that means the resource did not
|
||||
// set its own data pointer. So, we need to set it back to nullptr.
|
||||
if (res->data == mount->data) {
|
||||
@@ -261,6 +277,25 @@ int VFS::Read(mjResource* resource, const void** buffer) {
|
||||
return kFailedToRead;
|
||||
}
|
||||
|
||||
mjtSize VFS::Write(mjResource* resource, const void* buffer, mjtSize nbytes) {
|
||||
if (resource) {
|
||||
if (!resource->provider) {
|
||||
const mjResource* mount = FindMount(resource->name);
|
||||
if (mount) {
|
||||
resource->provider = mount->provider;
|
||||
// Smuggle mount data if not already set.
|
||||
if (!resource->data) {
|
||||
resource->data = mount->data;
|
||||
}
|
||||
}
|
||||
}
|
||||
if (resource->provider && resource->provider->write) {
|
||||
return resource->provider->write(resource, buffer, nbytes);
|
||||
}
|
||||
}
|
||||
return -1;
|
||||
}
|
||||
|
||||
VFS::ResourcePtr VFS::CreateResource(std::string_view name,
|
||||
const mjpResourceProvider* provider) {
|
||||
mjResource* res = new mjResource();
|
||||
|
||||
+9
-3
@@ -72,15 +72,21 @@ class VFS {
|
||||
// Opens a mjResource for the given path, or nullptr on error. If successful,
|
||||
// will invoke the 'open' callback for the mjpResourceProvider associated with
|
||||
// the path/dir.
|
||||
mjResource* Open(const char* dir, const char* name);
|
||||
mjResource* Open(const char* dir, const char* name, char* error = nullptr,
|
||||
size_t nerror = 0);
|
||||
|
||||
// Sets `buffer` to the contents of the resource and returns the number of
|
||||
// bytes of the content. This is done by invoking the 'read' callback for the
|
||||
// mjpResourceProvider associated with the resource. Returns -1 on error.
|
||||
int Read(mjResource* resource, const void** buffer);
|
||||
|
||||
// Closes the resource by invoking the 'close' callback for the
|
||||
// mjpResourceProvider associated with the resource.
|
||||
// Writes the resource data. This is done by invoking the 'write' callback for
|
||||
// the mjpResourceProvider associated with the resource.
|
||||
// Returns bytes written, or -1 on error.
|
||||
mjtSize Write(mjResource* resource, const void* buffer, mjtSize nbytes);
|
||||
|
||||
// Closes the resource by invoking the 'close' callback for
|
||||
// the mjpResourceProvider associated with the resource.
|
||||
Status Close(mjResource* resource);
|
||||
|
||||
// Mounts a ResourceProvider at the given path. All subsequent operations
|
||||
|
||||
+7
-4
@@ -186,10 +186,13 @@ int mj_saveXML(const mjSpec* s, const char* filename, char* error, int error_sz)
|
||||
return -1;
|
||||
}
|
||||
|
||||
std::ofstream file;
|
||||
file.open(filename);
|
||||
file << result;
|
||||
file.close();
|
||||
mjtSize written = mju_writeResource(filename, result.data(), result.size(), NULL, error, error_sz);
|
||||
if (written != result.size()) {
|
||||
if (error && error_sz > 0 && error[0] == '\0') {
|
||||
std::snprintf(error, error_sz, "Error writing XML file '%s'", filename);
|
||||
}
|
||||
return -1;
|
||||
}
|
||||
return 0;
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user