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:
committed by
Copybara-Service
parent
5ae6b5fe31
commit
61f13cd8f1
@@ -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);
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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
@@ -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) \
|
||||
|
||||
Reference in New Issue
Block a user