Add custom Emscripten binding for mj_saveModel.

Also adds support to unpack nullable string args

PiperOrigin-RevId: 859702186
Change-Id: Ia9e10f447ecf17b1dd74dd64d988790107a438a6
This commit is contained in:
Google DeepMind
2026-01-22 11:35:23 -08:00
committed by Copybara-Service
parent 5ae6b5fe31
commit 61f13cd8f1
5 changed files with 115 additions and 5 deletions
+16
View File
@@ -50,6 +50,7 @@ using emscripten::return_value_policy::take_ownership;
EMSCRIPTEN_DECLARE_VAL_TYPE(NumberOrString);
EMSCRIPTEN_DECLARE_VAL_TYPE(NumberArray);
EMSCRIPTEN_DECLARE_VAL_TYPE(String);
EMSCRIPTEN_DECLARE_VAL_TYPE(StringOrNull);
// Macro to define accessors for different MuJoCo object types within mjModel.
// Each line calls X_ACCESSOR with the following arguments:
@@ -8010,6 +8011,12 @@ std::unique_ptr<MjModel> mj_loadXML_wrapper(std::string filename) {
return std::unique_ptr<MjModel>(new MjModel(model));
}
void mj_saveModel_wrapper(const MjModel& m, const StringOrNull& filename, const val& buffer) {
UNPACK_NULLABLE_STRING(filename);
UNPACK_NULLABLE_VALUE(uint8_t, buffer);
mj_saveModel(m.get(), filename_.data(), buffer_.data(), static_cast<int>(buffer_.size()));
}
std::unique_ptr<MjSpec> parseXMLString_wrapper(const std::string &xml) {
char error[1000];
mjSpec *ptr = mj_parseXMLString(xml.c_str(), nullptr, error, sizeof(error));
@@ -12945,6 +12952,7 @@ EMSCRIPTEN_BINDINGS(mujoco_bindings) {
function("mjv_updateSkin", &mjv_updateSkin_wrapper);
function("parseXMLString", &parseXMLString_wrapper, take_ownership());
function("error", &error_wrapper);
function("mj_saveModel", &mj_saveModel_wrapper);
function("mj_saveLastXML", &mj_saveLastXML_wrapper);
function("mj_setLengthRange", &mj_setLengthRange_wrapper);
// mj_compile is bound using two overloads to handle the optional MjVFS argument,
@@ -12973,6 +12981,13 @@ EMSCRIPTEN_BINDINGS(mujoco_bindings) {
.function("GetElementCount", &WasmBuffer<int>::GetElementCount)
.function("GetView", &WasmBuffer<int>::GetView);
emscripten::class_<WasmBuffer<uint8_t>>("Uint8Buffer")
.constructor<int>()
.class_function("FromArray", &WasmBuffer<uint8_t>::FromArray)
.function("GetPointer", &WasmBuffer<uint8_t>::GetPointer)
.function("GetElementCount", &WasmBuffer<uint8_t>::GetElementCount)
.function("GetView", &WasmBuffer<uint8_t>::GetView);
emscripten::register_vector<std::string>("mjStringVec");
emscripten::register_vector<int>("mjIntVec");
emscripten::register_vector<mjIntVec>("mjIntVecVec");
@@ -12994,6 +13009,7 @@ EMSCRIPTEN_BINDINGS(mujoco_bindings) {
emscripten::register_type<NumberOrString>("number|string");
emscripten::register_type<NumberArray>("number[]");
emscripten::register_type<String>("string");
emscripten::register_type<StringOrNull>("string|null");
emscripten::constant("mjMAXCONPAIR", mjMAXCONPAIR);
emscripten::constant("mjMAXIMP", mjMAXIMP);
+1 -1
View File
@@ -142,7 +142,6 @@ _SKIPPED_MEMORY_FUNCTIONS: tuple[str, ...] = (
"mj_loadModel",
"mj_loadModelBuffer",
"mj_markStack",
"mj_saveModel",
"mj_stackAllocByte",
"mj_stackAllocInt",
"mj_stackAllocNum",
@@ -229,6 +228,7 @@ MANUAL_WRAPPER_FUNCTIONS: tuple[str, ...] = (
# go/keep-sorted start
"mj_compile",
"mj_saveLastXML",
"mj_saveModel",
"mj_setLengthRange",
"mju_error",
# go/keep-sorted end
+16
View File
@@ -50,6 +50,7 @@ using emscripten::return_value_policy::take_ownership;
EMSCRIPTEN_DECLARE_VAL_TYPE(NumberOrString);
EMSCRIPTEN_DECLARE_VAL_TYPE(NumberArray);
EMSCRIPTEN_DECLARE_VAL_TYPE(String);
EMSCRIPTEN_DECLARE_VAL_TYPE(StringOrNull);
// Macro to define accessors for different MuJoCo object types within mjModel.
// Each line calls X_ACCESSOR with the following arguments:
@@ -528,6 +529,12 @@ std::unique_ptr<MjModel> mj_loadXML_wrapper(std::string filename) {
return std::unique_ptr<MjModel>(new MjModel(model));
}
void mj_saveModel_wrapper(const MjModel& m, const StringOrNull& filename, const val& buffer) {
UNPACK_NULLABLE_STRING(filename);
UNPACK_NULLABLE_VALUE(uint8_t, buffer);
mj_saveModel(m.get(), filename_.data(), buffer_.data(), static_cast<int>(buffer_.size()));
}
std::unique_ptr<MjSpec> parseXMLString_wrapper(const std::string &xml) {
char error[1000];
mjSpec *ptr = mj_parseXMLString(xml.c_str(), nullptr, error, sizeof(error));
@@ -607,6 +614,7 @@ EMSCRIPTEN_BINDINGS(mujoco_bindings) {
// {{ FUNCTION_BINDINGS }}
function("parseXMLString", &parseXMLString_wrapper, take_ownership());
function("error", &error_wrapper);
function("mj_saveModel", &mj_saveModel_wrapper);
function("mj_saveLastXML", &mj_saveLastXML_wrapper);
function("mj_setLengthRange", &mj_setLengthRange_wrapper);
// mj_compile is bound using two overloads to handle the optional MjVFS argument,
@@ -635,6 +643,13 @@ EMSCRIPTEN_BINDINGS(mujoco_bindings) {
.function("GetElementCount", &WasmBuffer<int>::GetElementCount)
.function("GetView", &WasmBuffer<int>::GetView);
emscripten::class_<WasmBuffer<uint8_t>>("Uint8Buffer")
.constructor<int>()
.class_function("FromArray", &WasmBuffer<uint8_t>::FromArray)
.function("GetPointer", &WasmBuffer<uint8_t>::GetPointer)
.function("GetElementCount", &WasmBuffer<uint8_t>::GetElementCount)
.function("GetView", &WasmBuffer<uint8_t>::GetView);
emscripten::register_vector<std::string>("mjStringVec");
emscripten::register_vector<int>("mjIntVec");
emscripten::register_vector<mjIntVec>("mjIntVecVec");
@@ -656,6 +671,7 @@ EMSCRIPTEN_BINDINGS(mujoco_bindings) {
emscripten::register_type<NumberOrString>("number|string");
emscripten::register_type<NumberArray>("number[]");
emscripten::register_type<String>("string");
emscripten::register_type<StringOrNull>("string|null");
emscripten::constant("mjMAXCONPAIR", mjMAXCONPAIR);
emscripten::constant("mjMAXIMP", mjMAXIMP);
+37 -1
View File
@@ -17,7 +17,7 @@ import 'jasmine';
import {MainModule, MjContact, MjContactVec, MjData, MjLROpt, MjModel,
MjOption, MjsGeom, MjSolverStat, MjSpec, MjStatistic, MjTimerStat, MjvCamera,
MjvFigure, MjvGeom, MjvGLCamera, MjvLight, MjvOption, MjvPerturb, MjvScene,
MjWarningStat, MjVFS} from '../dist/mujoco_wasm.js';
MjWarningStat, MjVFS, Uint8Buffer} from '../dist/mujoco_wasm.js';
import loadMujoco from '../dist/mujoco_wasm.js'
@@ -2140,4 +2140,40 @@ describe('MuJoCo WASM Bindings', () => {
}
});
});
it('should save model to buffer', () => {
assertExists(model);
const modelSize = mujoco.mj_sizeModel(model);
const buffer = new mujoco.Uint8Buffer(Number(modelSize));
try {
mujoco.mj_saveModel(model, null, buffer);
// Test the first 3 byte values of the buffer.
expect(buffer.GetView()[0]).toBe(49);
expect(buffer.GetView()[1]).toBe(212);
expect(buffer.GetView()[2]).toBe(0);
} finally {
buffer.delete();
}
});
it('should save model to file', () => {
assertExists(model);
const filename = '/tmp/test.mjb';
const modelSize = mujoco.mj_sizeModel(model);
const buffer = new mujoco.Uint8Buffer(Number(modelSize));
try {
mujoco.mj_saveModel(model, filename, null);
const fileContent =
(mujoco as any).FS.readFile(filename, {encoding: 'binary'});
mujoco.mj_saveModel(model, null, buffer);
const bufferContent = buffer.GetView();
expect(fileContent.length).toBe(bufferContent.length);
expect(fileContent).toEqual(bufferContent);
} finally {
buffer.delete();
unlinkXMLFile(filename);
}
});
});
+45 -3
View File
@@ -97,7 +97,7 @@ class WasmBuffer {
template <typename T>
class UnpackedParam {
// The C++ representation of the parameter data
std::variant<std::monostate, std::vector<T>, std::span<T>> data_;
std::variant<std::monostate, std::vector<T>, std::span<T>, std::string> data_;
// Printable representations of the param and function name used for errors
const char* repr_;
@@ -112,7 +112,11 @@ class UnpackedParam {
UnpackedParam(T* data, std::size_t count, const char* repr, const char* func)
: data_(std::span<T>(data, count)), repr_(repr), func_(func) {}
UnpackedParam(std::string&& str, const char* repr, const char* func)
: data_(std::move(str)), repr_(repr), func_(func) {}
// Returns true and raises an error if the val is null or undefined.
// This function should never be called when unpacking nullable values.
static bool ErrorOnNullOrUndefined(const emscripten::val& p,
const char* func,
const char* expected_type) {
@@ -161,6 +165,24 @@ class UnpackedParam {
return UnpackedParam<T>(convertJSArrayToNumberVector<T>(p), repr, func);
}
// Create from a nullable Javascript string. Call via UNPACK_NULLABLE_STRING.
static UnpackedParam<T> FromNullableString(const emscripten::val& p,
const char* repr,
const char* func) {
if (IsNullOrUndefined(p)) {
return UnpackedParam<T>(repr, func);
}
if (!p.isString()) {
mju_error(
"[%s] Invalid argument. Expected a string for %s.",
StripWrapperSuffix(func).c_str(), repr);
return UnpackedParam<T>(repr, func);
}
static_assert(std::is_same_v<T, char>,
"UNPACK_NULLABLE_STRING requires UnpackedParam<char>");
return UnpackedParam<T>(p.as<std::string>(), repr, func);
}
// Creates an UnpackedParam from a Javascript a TypedArray or a WasmBuffer.
// Call via UNPACK_VALUE.
static UnpackedParam<T> FromValue(const emscripten::val& p, const char* repr,
@@ -200,14 +222,16 @@ class UnpackedParam {
// Returns the name of the function the parameter is used in.
std::string func() const { return StripWrapperSuffix(func_); }
// Returns the size of the parameter. Returns 0 if the parameter is null.
// Returns the size of the parameter. Returns 0 if the parameter is null. For
// strings, returns the length of the string.
std::size_t size() const {
if (std::holds_alternative<std::vector<T>>(data_)) {
return std::get<std::vector<T>>(data_).size();
} else if (std::holds_alternative<std::span<T>>(data_)) {
return std::get<std::span<T>>(data_).size();
} else if (std::holds_alternative<std::string>(data_)) {
return std::get<std::string>(data_).length();
}
mju_error("[%s] [%s] UnpackedParam is null", func().c_str(), repr());
return 0;
}
@@ -218,6 +242,11 @@ class UnpackedParam {
return std::get<std::vector<T>>(data_).data();
} else if (std::holds_alternative<std::span<T>>(data_)) {
return std::get<std::span<T>>(data_).data();
} else if (std::holds_alternative<std::string>(data_)) {
static_assert(std::is_same_v<T, char>,
"Cannot call data() on UnpackedParam with a string unless "
"T is char.");
return reinterpret_cast<const T*>(std::get<std::string>(data_).data());
}
return nullptr;
}
@@ -229,6 +258,14 @@ class UnpackedParam {
return std::get<std::vector<T>>(data_).data();
} else if (std::holds_alternative<std::span<T>>(data_)) {
return const_cast<T*>(std::get<std::span<T>>(data_).data());
} else if (std::holds_alternative<std::string>(data_)) {
if constexpr (std::is_same_v<T, char>) {
return reinterpret_cast<T*>(std::get<std::string>(data_).data());
} else {
mju_error(
"[%s] [%s] Cannot call data() on UnpackedParam<%s> holding a string",
func().c_str(), repr(), typeid(T).name());
}
}
return nullptr;
}
@@ -256,6 +293,11 @@ class UnpackedParam {
#define UNPACK_NULLABLE_ARRAY(T, p) \
UnpackedParam<T> p##_ = UnpackedParam<T>::FromNullableArray(p, #p, __func__)
#define UNPACK_NULLABLE_STRING(p) \
UnpackedParam<char> p##_ = UnpackedParam<char>::FromNullableString( \
p, #p, __func__ \
)
// Raises an error if x##_.size() is not equal to expr.
// Assumes UnpackedParam x##_ is defined.
#define CHECK_SIZE(x, expr) \