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
+11
View File
@@ -1639,6 +1639,17 @@ Close a resource; no-op if resource is NULL.
Set buffer to bytes read from the resource and return number of bytes in buffer;
return negative value if error.
.. _mju_writeResource:
`mju_writeResource <#mju_writeResource>`__
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
.. mujoco-include:: mju_writeResource
Write resource data via its resource provider, return bytes written or -1 on error.
*Nullable:* ``vfs``, ``error``
.. _mju_getResourceDir:
`mju_getResourceDir <#mju_getResourceDir>`__
+1
View File
@@ -20,6 +20,7 @@ General
- Fixed loading of ``.mjz`` archives in :ref:`simulate<saSimulate>`: the archive was unmounted before model compilation,
so assets contained in it failed to load. Failures in the ``mjz`` decoder now emit a warning with the underlying
error instead of the generic "could not decode content" message.
- Added support for resource writing via :ref:`mju_writeResource` and the ``write`` callback in :ref:`mjpResourceProvider`.
.. admonition:: Breaking API changes
:class: attention
+3
View File
@@ -1218,6 +1218,7 @@ typedef struct mjpResourceProvider {
mjfMountResource mount; // mounting callback (optional)
mjfUnmountResource unmount; // unmounting callback (optional)
mjfResourceModified modified; // resource modified callback (optional)
mjfWriteResource write; // writing callback (optional)
void* data; // opaque data pointer (resource invariant)
} mjpResourceProvider;
typedef struct mjpDecoder {
@@ -3684,6 +3685,8 @@ mjResource* mju_openResource(const char* dir, const char* name,
const mjVFS* vfs, char* error, size_t nerror);
void mju_closeResource(mjResource* resource);
int mju_readResource(mjResource* resource, const void** buffer);
mjtSize mju_writeResource(const char* name, const void* buffer, mjtSize nbytes,
const mjVFS* vfs, char* error, size_t nerror);
void mju_getResourceDir(mjResource* resource, const char** dir, int* ndir);
int mju_isModifiedResource(const mjResource* resource, const char* timestamp);
mjSpec* mju_decodeResource(mjResource* resource, const char* content_type,
+5
View File
@@ -58,6 +58,10 @@ typedef int (*mjfUnmountResource)(mjResource* resource);
// returns < 0 if the resource is older than the given timestamp
typedef int (*mjfResourceModified)(const mjResource* resource, const char* timestamp);
// callback for writing bytes to a resource
// return number of bytes written, return -1 if error
typedef mjtSize (*mjfWriteResource)(mjResource* resource, const void* buffer, mjtSize nbytes);
// struct describing a single resource provider
typedef struct mjpResourceProvider {
const char* prefix; // prefix for match against a resource name
@@ -67,6 +71,7 @@ typedef struct mjpResourceProvider {
mjfMountResource mount; // mounting callback (optional)
mjfUnmountResource unmount; // unmounting callback (optional)
mjfResourceModified modified; // resource modified callback (optional)
mjfWriteResource write; // writing callback (optional)
void* data; // opaque data pointer (resource invariant)
} mjpResourceProvider;
+5
View File
@@ -1588,6 +1588,11 @@ MJAPI void mju_closeResource(mjResource* resource);
// return negative value if error.
MJAPI int mju_readResource(mjResource* resource, const void** buffer);
// Write resource data via its resource provider, return bytes written or -1 on error.
// Nullable: vfs, error
MJAPI mjtSize mju_writeResource(const char* name, const void* buffer, mjtSize nbytes,
const mjVFS* vfs, char* error, size_t nerror);
// For a resource with a name partitioned as {dir}{filename}, get the dir and ndir pointers.
MJAPI void mju_getResourceDir(mjResource* resource, const char** dir, int* ndir);
@@ -265,6 +265,7 @@ _OMITTED_FUNCTIONS = [
'mju_openResource',
'mju_closeResource',
'mju_readResource',
'mju_writeResource',
'mju_getResourceDir',
'mju_isModifiedResource',
'mj_compile',
+42
View File
@@ -9929,6 +9929,48 @@ FUNCTIONS: Mapping[str, FunctionDecl] = dict([
),
doc='Set buffer to bytes read from the resource and return number of bytes in buffer; return negative value if error.', # pylint: disable=line-too-long
)),
('mju_writeResource',
FunctionDecl(
name='mju_writeResource',
return_type=ValueType(name='mjtSize'),
parameters=(
FunctionParameterDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='char', is_const=True),
),
),
FunctionParameterDecl(
name='buffer',
type=PointerType(
inner_type=ValueType(name='void', is_const=True),
),
),
FunctionParameterDecl(
name='nbytes',
type=ValueType(name='mjtSize'),
),
FunctionParameterDecl(
name='vfs',
type=PointerType(
inner_type=ValueType(name='mjVFS', is_const=True),
),
nullable=True,
),
FunctionParameterDecl(
name='error',
type=PointerType(
inner_type=ValueType(name='char'),
),
nullable=True,
),
FunctionParameterDecl(
name='nerror',
type=ValueType(name='size_t'),
),
),
doc='Write resource data via its resource provider, return bytes written or -1 on error.', # pylint: disable=line-too-long
)),
('mju_getResourceDir',
FunctionDecl(
name='mju_getResourceDir',
+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;
}
+2 -2
View File
@@ -48,7 +48,7 @@ mjtSize FakeEncode(const mjSpec* s, const mjModel* m, const mjVFS* vfs,
void CloseResource(mjResource* resource) {
delete static_cast<FakeEncoderOutput*>(resource->data);
delete resource;
resource->data = nullptr;
}
mjpEncoder FakeEncoder() {
@@ -125,7 +125,7 @@ TEST_F(EncoderPluginTest, EncodeModel) {
EXPECT_EQ(output->njnt, 0);
EXPECT_STREQ(output->resource_name, "output.fakeformat");
delete output;
found->close_resource(&resource);
mj_deleteModel(model);
mj_deleteSpec(spec);
}
+104 -1
View File
@@ -14,8 +14,11 @@
// Tests for user/user_resource.cc
#include "src/user/user_resource.h"
#include <array>
#include <cstdint>
#include <cstdio>
#include <cstring>
#include <ctime>
#include <string>
@@ -29,7 +32,6 @@
#include <mujoco/mujoco.h>
#include "src/engine/engine_plugin.h"
#include "src/engine/engine_util_misc.h"
#include "src/user/user_resource.h"
#include "test/fixture.h"
namespace mujoco {
@@ -441,5 +443,106 @@ TEST_F(ResourceTest, OSFilesystemTimestamps) {
mju_closeResource(resource);
}
// ===================== Write Resource Tests =====================
struct WriteBuffer {
std::vector<uint8_t> data;
};
// static storage for the last written buffer (for round-trip testing)
static std::vector<uint8_t> last_written;
mjtSize write_capture(mjResource* resource, const void* buffer, mjtSize nbytes) {
last_written.clear();
const uint8_t* bytes = static_cast<const uint8_t*>(buffer);
last_written.insert(last_written.end(), bytes, bytes + nbytes);
return nbytes;
}
// open callback for reading back captured data
int open_read_capture(mjResource* resource) { return 1; }
int read_capture(mjResource* resource, const void** buffer) {
*buffer = last_written.data();
return static_cast<int>(last_written.size());
}
void close_read_capture(mjResource* resource) {}
TEST_F(ResourceTest, WriteResourceWithProvider) {
// Register a provider with write callbacks
mjpResourceProvider provider = {
.prefix = "wrtest",
.open = open_read_capture,
.read = read_capture,
.close = close_read_capture,
.write = write_capture,
};
mjp_registerResourceProvider(&provider);
// Write some data
const char* data = "Hello, Resource Writer!";
mjtSize nbytes = static_cast<mjtSize>(std::strlen(data));
mjtSize written = mju_writeResource("wrtest:myfile.txt", data, nbytes, nullptr, nullptr, 0);
EXPECT_EQ(written, nbytes);
// Read it back via the same provider
char error[256] = {0};
mjResource* resource =
mju_openResource("", "wrtest:myfile.txt", nullptr, error, sizeof(error));
ASSERT_THAT(resource, NotNull());
const void* buf = nullptr;
int read_bytes = mju_readResource(resource, &buf);
EXPECT_EQ(read_bytes, nbytes);
EXPECT_EQ(std::memcmp(buf, data, nbytes), 0);
mju_closeResource(resource);
}
TEST_F(ResourceTest, WriteResourceDefaultPosix) {
// Write to a temp file using the default POSIX provider (no prefix)
std::string tmpfile = testing::TempDir() + "/mj_write_test.bin";
const char* data = "MuJoCo resource write test";
mjtSize nbytes = static_cast<mjtSize>(std::strlen(data));
mjtSize written = mju_writeResource(tmpfile.c_str(), data, nbytes, nullptr, nullptr, 0);
EXPECT_EQ(written, nbytes);
// Read it back via mju_openResource (default POSIX provider)
char error[256] = {0};
mjResource* resource =
mju_openResource("", tmpfile.c_str(), nullptr, error, sizeof(error));
ASSERT_THAT(resource, NotNull());
const void* buf = nullptr;
int read_bytes = mju_readResource(resource, &buf);
EXPECT_EQ(read_bytes, nbytes);
EXPECT_EQ(std::memcmp(buf, data, nbytes), 0);
mju_closeResource(resource);
// Clean up
std::remove(tmpfile.c_str());
}
TEST_F(ResourceTest, WriteResourceNoWriteCallback) {
// Register a read-only provider (no write callback)
mjpResourceProvider provider = {
.prefix = "rdonly",
.open = open_nop,
.read = read_nop,
.close = close_nop,
};
mjp_registerResourceProvider(&provider);
const char* data = "data";
mjtSize written = mju_writeResource("rdonly:somefile", data, 4, nullptr, nullptr, 0);
EXPECT_EQ(written, -1);
}
} // namespace
} // namespace mujoco
+1
View File
@@ -193,6 +193,7 @@ _SKIPPED_UTILITY_FUNCTIONS: tuple[str, ...] = (
"mju_isModifiedResource",
"mju_openResource",
"mju_readResource",
"mju_writeResource",
# go/keep-sorted end
)