From 61f13cd8f1e42bc64ea9b08b670c94a64481845c Mon Sep 17 00:00:00 2001 From: Google DeepMind Date: Thu, 22 Jan 2026 11:35:23 -0800 Subject: [PATCH] Add custom Emscripten binding for `mj_saveModel`. Also adds support to unpack nullable string args PiperOrigin-RevId: 859702186 Change-Id: Ia9e10f447ecf17b1dd74dd64d988790107a438a6 --- wasm/codegen/generated/bindings.cc | 16 ++++++++++ wasm/codegen/generators/constants.py | 2 +- wasm/codegen/templates/bindings.cc | 16 ++++++++++ wasm/tests/bindings_test.ts | 38 +++++++++++++++++++++- wasm/unpack.h | 48 ++++++++++++++++++++++++++-- 5 files changed, 115 insertions(+), 5 deletions(-) diff --git a/wasm/codegen/generated/bindings.cc b/wasm/codegen/generated/bindings.cc index 4ad90672..4b69b9e4 100644 --- a/wasm/codegen/generated/bindings.cc +++ b/wasm/codegen/generated/bindings.cc @@ -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 mj_loadXML_wrapper(std::string filename) { return std::unique_ptr(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(buffer_.size())); +} + std::unique_ptr 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::GetElementCount) .function("GetView", &WasmBuffer::GetView); + emscripten::class_>("Uint8Buffer") + .constructor() + .class_function("FromArray", &WasmBuffer::FromArray) + .function("GetPointer", &WasmBuffer::GetPointer) + .function("GetElementCount", &WasmBuffer::GetElementCount) + .function("GetView", &WasmBuffer::GetView); + emscripten::register_vector("mjStringVec"); emscripten::register_vector("mjIntVec"); emscripten::register_vector("mjIntVecVec"); @@ -12994,6 +13009,7 @@ EMSCRIPTEN_BINDINGS(mujoco_bindings) { emscripten::register_type("number|string"); emscripten::register_type("number[]"); emscripten::register_type("string"); + emscripten::register_type("string|null"); emscripten::constant("mjMAXCONPAIR", mjMAXCONPAIR); emscripten::constant("mjMAXIMP", mjMAXIMP); diff --git a/wasm/codegen/generators/constants.py b/wasm/codegen/generators/constants.py index 0fe8b6e2..5046ab1e 100644 --- a/wasm/codegen/generators/constants.py +++ b/wasm/codegen/generators/constants.py @@ -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 diff --git a/wasm/codegen/templates/bindings.cc b/wasm/codegen/templates/bindings.cc index f7b4df62..88a640ff 100644 --- a/wasm/codegen/templates/bindings.cc +++ b/wasm/codegen/templates/bindings.cc @@ -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 mj_loadXML_wrapper(std::string filename) { return std::unique_ptr(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(buffer_.size())); +} + std::unique_ptr 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::GetElementCount) .function("GetView", &WasmBuffer::GetView); + emscripten::class_>("Uint8Buffer") + .constructor() + .class_function("FromArray", &WasmBuffer::FromArray) + .function("GetPointer", &WasmBuffer::GetPointer) + .function("GetElementCount", &WasmBuffer::GetElementCount) + .function("GetView", &WasmBuffer::GetView); + emscripten::register_vector("mjStringVec"); emscripten::register_vector("mjIntVec"); emscripten::register_vector("mjIntVecVec"); @@ -656,6 +671,7 @@ EMSCRIPTEN_BINDINGS(mujoco_bindings) { emscripten::register_type("number|string"); emscripten::register_type("number[]"); emscripten::register_type("string"); + emscripten::register_type("string|null"); emscripten::constant("mjMAXCONPAIR", mjMAXCONPAIR); emscripten::constant("mjMAXIMP", mjMAXIMP); diff --git a/wasm/tests/bindings_test.ts b/wasm/tests/bindings_test.ts index 5501fbe2..c46d79a5 100644 --- a/wasm/tests/bindings_test.ts +++ b/wasm/tests/bindings_test.ts @@ -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); + } + }); }); diff --git a/wasm/unpack.h b/wasm/unpack.h index 9942a773..b4e57bf7 100644 --- a/wasm/unpack.h +++ b/wasm/unpack.h @@ -97,7 +97,7 @@ class WasmBuffer { template class UnpackedParam { // The C++ representation of the parameter data - std::variant, std::span> data_; + std::variant, std::span, 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(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(convertJSArrayToNumberVector(p), repr, func); } + // Create from a nullable Javascript string. Call via UNPACK_NULLABLE_STRING. + static UnpackedParam FromNullableString(const emscripten::val& p, + const char* repr, + const char* func) { + if (IsNullOrUndefined(p)) { + return UnpackedParam(repr, func); + } + if (!p.isString()) { + mju_error( + "[%s] Invalid argument. Expected a string for %s.", + StripWrapperSuffix(func).c_str(), repr); + return UnpackedParam(repr, func); + } + static_assert(std::is_same_v, + "UNPACK_NULLABLE_STRING requires UnpackedParam"); + return UnpackedParam(p.as(), repr, func); + } + // Creates an UnpackedParam from a Javascript a TypedArray or a WasmBuffer. // Call via UNPACK_VALUE. static UnpackedParam 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>(data_)) { return std::get>(data_).size(); } else if (std::holds_alternative>(data_)) { return std::get>(data_).size(); + } else if (std::holds_alternative(data_)) { + return std::get(data_).length(); } - mju_error("[%s] [%s] UnpackedParam is null", func().c_str(), repr()); return 0; } @@ -218,6 +242,11 @@ class UnpackedParam { return std::get>(data_).data(); } else if (std::holds_alternative>(data_)) { return std::get>(data_).data(); + } else if (std::holds_alternative(data_)) { + static_assert(std::is_same_v, + "Cannot call data() on UnpackedParam with a string unless " + "T is char."); + return reinterpret_cast(std::get(data_).data()); } return nullptr; } @@ -229,6 +258,14 @@ class UnpackedParam { return std::get>(data_).data(); } else if (std::holds_alternative>(data_)) { return const_cast(std::get>(data_).data()); + } else if (std::holds_alternative(data_)) { + if constexpr (std::is_same_v) { + return reinterpret_cast(std::get(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 p##_ = UnpackedParam::FromNullableArray(p, #p, __func__) +#define UNPACK_NULLABLE_STRING(p) \ + UnpackedParam p##_ = UnpackedParam::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) \