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:
Sam Haves
2026-07-09 10:47:01 -07:00
committed by Copybara-Service
parent bdeb7e7c4d
commit dc7581acfa
18 changed files with 324 additions and 83 deletions
+25 -38
View File
@@ -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);
}
}
+5 -5
View File
@@ -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
View File
@@ -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
View File
@@ -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;
+3
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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;
}