From c996f288e15f14306f0ad9d5ae332325f677af95 Mon Sep 17 00:00:00 2001 From: Google DeepMind Date: Wed, 14 Jan 2026 07:30:43 -0800 Subject: [PATCH] Add code generation for mjModel named accessors. Most of it is done leveraging existing X Macros (indexers.h, indexer_xmacro.h) used for python bindings already. PiperOrigin-RevId: 856205245 Change-Id: Ic72443999065040af3042bbf3f9f68c5b70d309d --- wasm/codegen/generated/bindings.cc | 155 ++++++++++++++++ wasm/codegen/generators/structs.py | 39 ++++- wasm/codegen/templates/bindings.cc | 127 ++++++++++++++ wasm/tests/bindings_test.ts | 272 ++++++++++++++++++++++++++++- 4 files changed, 588 insertions(+), 5 deletions(-) diff --git a/wasm/codegen/generated/bindings.cc b/wasm/codegen/generated/bindings.cc index 164437c5..def6244a 100644 --- a/wasm/codegen/generated/bindings.cc +++ b/wasm/codegen/generated/bindings.cc @@ -28,6 +28,7 @@ #include #include // NOLINT #include // NOLINT +#include #include #include @@ -36,6 +37,7 @@ #include #include "engine/engine_util_errmem.h" #include "wasm/unpack.h" +#include "python/mujoco/indexer_xmacro.h" namespace mujoco::wasm { @@ -45,9 +47,39 @@ using emscripten::val; using emscripten::return_value_policy::reference; using emscripten::return_value_policy::take_ownership; +EMSCRIPTEN_DECLARE_VAL_TYPE(NumberOrString); EMSCRIPTEN_DECLARE_VAL_TYPE(NumberArray); EMSCRIPTEN_DECLARE_VAL_TYPE(String); +// Macro to define accessors for different MuJoCo object types within mjModel. +// Each line calls X_ACCESSOR with the following arguments: +// 1. NAME: The object type name in uppercase (e.g., ACTUATOR). +// 2. Name: The object type name in CamelCase (e.g., Actuator). +// 3. OBJTYPE: The corresponding mjOBJ_* enum value (e.g., mjOBJ_ACTUATOR). +// 4. field: The name of the array field in mjModel (e.g., actuator). +// 5. nfield: The name of the count field in mjModel (e.g., nu). +#define MJMODEL_ACCESSORS \ + X_ACCESSOR( ACTUATOR, Actuator, mjOBJ_ACTUATOR, actuator, nu ) \ + X_ACCESSOR( BODY, Body, mjOBJ_BODY, body, nbody ) \ + X_ACCESSOR( CAMERA, Camera, mjOBJ_CAMERA, cam, ncam ) \ + X_ACCESSOR( EQUALITY, Equality, mjOBJ_EQUALITY, eq, neq ) \ + X_ACCESSOR( EXCLUDE, Exclude, mjOBJ_EXCLUDE, exclude, nexclude ) \ + X_ACCESSOR( GEOM, Geom, mjOBJ_GEOM, geom, ngeom ) \ + X_ACCESSOR( HFIELD, Hfield, mjOBJ_HFIELD, hfield, nhfield ) \ + X_ACCESSOR( JOINT, Joint, mjOBJ_JOINT, jnt, njnt ) \ + X_ACCESSOR( LIGHT, Light, mjOBJ_LIGHT, light, nlight ) \ + X_ACCESSOR( MATERIAL, Material, mjOBJ_MATERIAL, mat, nmat ) \ + X_ACCESSOR( MESH, Mesh, mjOBJ_MESH, mesh, nmesh ) \ + X_ACCESSOR( NUMERIC, Numeric, mjOBJ_NUMERIC, numeric, nnumeric ) \ + X_ACCESSOR( PAIR, Pair, mjOBJ_PAIR, pair, npair ) \ + X_ACCESSOR( SENSOR, Sensor, mjOBJ_SENSOR, sensor, nsensor ) \ + X_ACCESSOR( SITE, Site, mjOBJ_SITE, site, nsite ) \ + X_ACCESSOR( SKIN, Skin, mjOBJ_SKIN, skin, nskin ) \ + X_ACCESSOR( TENDON, Tendon, mjOBJ_TENDON, tendon, ntendon ) \ + X_ACCESSOR( TEXTURE, Texture, mjOBJ_TEXTURE, tex, ntex ) \ + X_ACCESSOR( TUPLE, Tuple, mjOBJ_TUPLE, tuple, ntuple ) \ + X_ACCESSOR( KEYFRAME, Keyframe, mjOBJ_KEY, key, nkey ) + // Raises an error if the given val is null or undefined. // A macro is used so that the error contains the name of the variable. // TODO(matijak): Remove this when we can handle strings using UNPACK_STRING? @@ -117,6 +149,84 @@ using mjVisualQuality = decltype(::mjVisual::quality); using mjVisualRgba = decltype(::mjVisual::rgba); using mjVisualScale = decltype(::mjVisual::scale); +#undef MJ_M +#define MJ_M(n) model_->n +// The X macro expands to a member function within an MjModel...Accessor struct. +// This function returns an emscripten::val, typically a typed memory view, +// providing access to array data within the underlying mjModel. The size and +// offset of the memory view are determined by the arguments: +// - type: The C++ type of the array elements (e.g., mjtNum, int). +// - prefix: The prefix of the field name in mjModel (e.g., jnt_, geom_). +// - var: The name of the field being accessed (e.g., qposadr, size). +// - dim0: Used to determine special indexing logic for dynamically sized arrays +// like joints (nq, nv), hfields, textures, numerics, and tuples. +// - dim1: The fixed dimension of the array if not dynamically sized. If "1", +// a single element is returned. Otherwise, it's used as the stride +// for the typed memory view. +#define X(type, prefix, var, dim0, dim1) \ + auto var() const { \ + if constexpr (std::string_view(#dim0) == "nq") { \ + int start = model_->jnt_qposadr[id_]; \ + int end = (id_ < model_->njnt - 1) ? model_->jnt_qposadr[id_ + 1] : model_->nq; \ + return emscripten::val(emscripten::typed_memory_view(end - start, model_->prefix##var + start)); \ + } else if constexpr (std::string_view(#dim0) == "nv") { \ + int start = model_->jnt_dofadr[id_]; \ + int end = (id_ < model_->njnt - 1) ? model_->jnt_dofadr[id_ + 1] : model_->nv; \ + return emscripten::val(emscripten::typed_memory_view(end - start, model_->prefix##var + start)); \ + } else if constexpr (std::string_view(#dim0) == "nhfielddata") { \ + int start = model_->hfield_adr[id_]; \ + int count = model_->hfield_nrow[id_] * model_->hfield_ncol[id_]; \ + return emscripten::val(emscripten::typed_memory_view(count, model_->hfield_data + start)); \ + } else if constexpr (std::string_view(#dim0) == "ntexdata") { \ + int start = model_->tex_adr[id_]; \ + int count = model_->tex_height[id_] * model_->tex_width[id_] * model_->tex_nchannel[id_]; \ + return emscripten::val(emscripten::typed_memory_view(count, model_->tex_data + start)); \ + } else if constexpr (std::string_view(#dim0) == "nnumericdata") { \ + int start = model_->numeric_adr[id_]; \ + int count = model_->numeric_size[id_]; \ + return emscripten::val(emscripten::typed_memory_view(count, model_->numeric_data + start)); \ + } else if constexpr (std::string_view(#dim0) == "ntupledata") { \ + int start = model_->tuple_adr[id_]; \ + int count = model_->tuple_size[id_]; \ + return emscripten::val(emscripten::typed_memory_view(count, model_->prefix##var + start)); \ + } else { \ + if constexpr (std::string_view(#dim1) == "1") { \ + return model_->prefix##var[id_]; \ + } else { \ + return emscripten::val( \ + emscripten::typed_memory_view(dim1, model_->prefix##var + id_ * dim1)); \ + } \ + } \ + } + +// Expands to a struct definition for each object type in MJMODEL_ACCESSORS. +// Each struct, named `MjModel{Name}Accessor`, provides: +// - A constructor taking an `mjModel*` and an integer `id`. +// - An `id()` method to get the object's index. +// - A `name()` method to get the object's name using `mj_id2name`. +// - Member functions generated by the `MJMODEL_##NAME` macro, which in turn +// uses the `X` macro to define accessors for fields within `mjModel`. +#define X_ACCESSOR(NAME, Name, OBJTYPE, field, nfield) \ + struct MjModel##Name##Accessor { \ + MjModel##Name##Accessor(mjModel* model, int id) : model_(model), id_(id) {} \ + \ + int id() const { return id_; } \ + std::string name() const { \ + return mj_id2name(model_, OBJTYPE, id_); \ + } \ + \ + MJMODEL_##NAME \ + \ + private: \ + mjModel* model_; \ + int id_; \ + }; +MJMODEL_ACCESSORS +#undef X_ACCESSOR + +#undef X +#undef MJ_M + struct MjContact { ~MjContact(); MjContact(); @@ -5143,6 +5253,29 @@ struct MjModel { void set_signature(uint64_t value) { ptr_->signature = value; } + // Generates functions to return accessor classes. + #define X_ACCESSOR(NAME, Name, OBJTYPE, accessor_name, nfield) \ + MjModel##Name##Accessor accessor_name(const NumberOrString& val) const { \ + if (val.isString()) { \ + int id = mj_name2id(ptr_, OBJTYPE, val.as().c_str()); \ + if (id == -1) { \ + mju_error("Invalid name, MjModel." #accessor_name " not found"); \ + } \ + return MjModel##Name##Accessor(ptr_, id); \ + } else if (val.isNumber()) { \ + int id = val.as(); \ + if (id < 0 || id >= ptr_->nfield) { \ + mju_error_i("Invalid id %d for MjModel." #accessor_name, id); \ + } \ + return MjModel##Name##Accessor(ptr_, id); \ + } else { \ + mju_error(#accessor_name "() argument must be a string or number"); \ + return MjModel##Name##Accessor(nullptr, 0); \ + } \ + } + + MJMODEL_ACCESSORS + #undef X_ACCESSOR private: mjModel* ptr_; @@ -10784,6 +10917,20 @@ EMSCRIPTEN_BINDINGS(mujoco_bindings) { .value("mjWRAP_SPHERE", mjWRAP_SPHERE) .value("mjWRAP_CYLINDER", mjWRAP_CYLINDER); + // Bindings for the MjModel accessor classes. + #define X(type, prefix, var, dim0, dim1) .property(#var, &Accessor::var) + #define X_ACCESSOR(NAME, Name, OBJTYPE, field, nfield) \ + { \ + using Accessor = MjModel##Name##Accessor; \ + emscripten::class_("MjModel" #Name "Accessor") \ + .property("id", &Accessor::id) \ + .property("name", &Accessor::name) \ + MJMODEL_##NAME; \ + } + MJMODEL_ACCESSORS + #undef X + #undef X_ACCESSOR + emscripten::class_("MjContact") .constructor<>() .function("copy", &MjContact::copy, take_ownership()) @@ -11018,6 +11165,11 @@ EMSCRIPTEN_BINDINGS(mujoco_bindings) { emscripten::class_("MjModel") .class_function("mj_loadXML", &mj_loadXML_wrapper, take_ownership()) .constructor() + // Binds the functions on MjModel that return accessors. + #define X_ACCESSOR(NAME, Name, OBJTYPE, field_name, nfield) \ + .function(#field_name, &MjModel::field_name) + MJMODEL_ACCESSORS + #undef X_ACCESSOR .property("B_colind", &MjModel::B_colind) .property("B_rowadr", &MjModel::B_rowadr) .property("B_rownnz", &MjModel::B_rownnz) @@ -12795,6 +12947,9 @@ EMSCRIPTEN_BINDINGS(mujoco_bindings) { emscripten::register_vector("MjvGeomVec"); // register_type() improves type information (val is mapped to any by default) + // NumberOrString is used in functions returning accessors, allowing users to + // get an accessor by name (string) or id (number). + emscripten::register_type("number|string"); emscripten::register_type("number[]"); emscripten::register_type("string"); diff --git a/wasm/codegen/generators/structs.py b/wasm/codegen/generators/structs.py index acae9cec..9eafd731 100644 --- a/wasm/codegen/generators/structs.py +++ b/wasm/codegen/generators/structs.py @@ -20,8 +20,8 @@ import math from typing import Tuple, Union, cast from introspect import ast_nodes -from introspect import structs as introspect_structs from introspect import functions as introspect_functions +from introspect import structs as introspect_structs from wasm.codegen.generators import code_builder from wasm.codegen.generators import common @@ -376,6 +376,33 @@ def build_struct_header( for line in field.declaration.splitlines(): builder.line(line) + # accessors declarations + if w == "MjModel": + builder.line(""" + // Generates functions to return accessor classes. + #define X_ACCESSOR(NAME, Name, OBJTYPE, accessor_name, nfield) \\ + MjModel##Name##Accessor accessor_name(const NumberOrString& val) const { \\ + if (val.isString()) { \\ + int id = mj_name2id(ptr_, OBJTYPE, val.as().c_str()); \\ + if (id == -1) { \\ + mju_error("Invalid name, MjModel." #accessor_name " not found"); \\ + } \\ + return MjModel##Name##Accessor(ptr_, id); \\ + } else if (val.isNumber()) { \\ + int id = val.as(); \\ + if (id < 0 || id >= ptr_->nfield) { \\ + mju_error_i("Invalid id %d for MjModel." #accessor_name, id); \\ + } \\ + return MjModel##Name##Accessor(ptr_, id); \\ + } else { \\ + mju_error(#accessor_name "() argument must be a string or number"); \\ + return MjModel##Name##Accessor(nullptr, 0); \\ + } \\ + } + + MJMODEL_ACCESSORS + #undef X_ACCESSOR""".lstrip()) + # define private struct members builder.private() builder.line(f"{s}* ptr_;") @@ -481,11 +508,17 @@ def _build_struct_bindings( builder.line(".constructor()") builder.line(".constructor()") elif w == "MjModel": - w = common.wrapped_function_name( + fn = common.wrapped_function_name( introspect_functions.FUNCTIONS["mj_loadXML"] ) - builder.line(f'.class_function("mj_loadXML", &{w}, take_ownership())') + builder.line(f'.class_function("mj_loadXML", &{fn}, take_ownership())') builder.line(".constructor()") + builder.line(""" + // Binds the functions on MjModel that return accessors. + #define X_ACCESSOR(NAME, Name, OBJTYPE, field_name, nfield) \\ + .function(#field_name, &MjModel::field_name) + MJMODEL_ACCESSORS + #undef X_ACCESSOR""".lstrip()) elif w == "MjSpec": builder.line(".constructor()") elif w == "MjvScene": diff --git a/wasm/codegen/templates/bindings.cc b/wasm/codegen/templates/bindings.cc index 5e040fc7..f7b4df62 100644 --- a/wasm/codegen/templates/bindings.cc +++ b/wasm/codegen/templates/bindings.cc @@ -28,6 +28,7 @@ #include #include // NOLINT #include // NOLINT +#include #include #include @@ -36,6 +37,7 @@ #include #include "engine/engine_util_errmem.h" #include "wasm/unpack.h" +#include "python/mujoco/indexer_xmacro.h" namespace mujoco::wasm { @@ -45,9 +47,39 @@ using emscripten::val; using emscripten::return_value_policy::reference; using emscripten::return_value_policy::take_ownership; +EMSCRIPTEN_DECLARE_VAL_TYPE(NumberOrString); EMSCRIPTEN_DECLARE_VAL_TYPE(NumberArray); EMSCRIPTEN_DECLARE_VAL_TYPE(String); +// Macro to define accessors for different MuJoCo object types within mjModel. +// Each line calls X_ACCESSOR with the following arguments: +// 1. NAME: The object type name in uppercase (e.g., ACTUATOR). +// 2. Name: The object type name in CamelCase (e.g., Actuator). +// 3. OBJTYPE: The corresponding mjOBJ_* enum value (e.g., mjOBJ_ACTUATOR). +// 4. field: The name of the array field in mjModel (e.g., actuator). +// 5. nfield: The name of the count field in mjModel (e.g., nu). +#define MJMODEL_ACCESSORS \ + X_ACCESSOR( ACTUATOR, Actuator, mjOBJ_ACTUATOR, actuator, nu ) \ + X_ACCESSOR( BODY, Body, mjOBJ_BODY, body, nbody ) \ + X_ACCESSOR( CAMERA, Camera, mjOBJ_CAMERA, cam, ncam ) \ + X_ACCESSOR( EQUALITY, Equality, mjOBJ_EQUALITY, eq, neq ) \ + X_ACCESSOR( EXCLUDE, Exclude, mjOBJ_EXCLUDE, exclude, nexclude ) \ + X_ACCESSOR( GEOM, Geom, mjOBJ_GEOM, geom, ngeom ) \ + X_ACCESSOR( HFIELD, Hfield, mjOBJ_HFIELD, hfield, nhfield ) \ + X_ACCESSOR( JOINT, Joint, mjOBJ_JOINT, jnt, njnt ) \ + X_ACCESSOR( LIGHT, Light, mjOBJ_LIGHT, light, nlight ) \ + X_ACCESSOR( MATERIAL, Material, mjOBJ_MATERIAL, mat, nmat ) \ + X_ACCESSOR( MESH, Mesh, mjOBJ_MESH, mesh, nmesh ) \ + X_ACCESSOR( NUMERIC, Numeric, mjOBJ_NUMERIC, numeric, nnumeric ) \ + X_ACCESSOR( PAIR, Pair, mjOBJ_PAIR, pair, npair ) \ + X_ACCESSOR( SENSOR, Sensor, mjOBJ_SENSOR, sensor, nsensor ) \ + X_ACCESSOR( SITE, Site, mjOBJ_SITE, site, nsite ) \ + X_ACCESSOR( SKIN, Skin, mjOBJ_SKIN, skin, nskin ) \ + X_ACCESSOR( TENDON, Tendon, mjOBJ_TENDON, tendon, ntendon ) \ + X_ACCESSOR( TEXTURE, Texture, mjOBJ_TEXTURE, tex, ntex ) \ + X_ACCESSOR( TUPLE, Tuple, mjOBJ_TUPLE, tuple, ntuple ) \ + X_ACCESSOR( KEYFRAME, Keyframe, mjOBJ_KEY, key, nkey ) + // Raises an error if the given val is null or undefined. // A macro is used so that the error contains the name of the variable. // TODO(matijak): Remove this when we can handle strings using UNPACK_STRING? @@ -112,6 +144,84 @@ val get_mjRNDSTRING() { return MakeValArray3(mjRNDSTRING); } // {{ ANONYMOUS_STRUCT_TYPEDEFS }} +#undef MJ_M +#define MJ_M(n) model_->n +// The X macro expands to a member function within an MjModel...Accessor struct. +// This function returns an emscripten::val, typically a typed memory view, +// providing access to array data within the underlying mjModel. The size and +// offset of the memory view are determined by the arguments: +// - type: The C++ type of the array elements (e.g., mjtNum, int). +// - prefix: The prefix of the field name in mjModel (e.g., jnt_, geom_). +// - var: The name of the field being accessed (e.g., qposadr, size). +// - dim0: Used to determine special indexing logic for dynamically sized arrays +// like joints (nq, nv), hfields, textures, numerics, and tuples. +// - dim1: The fixed dimension of the array if not dynamically sized. If "1", +// a single element is returned. Otherwise, it's used as the stride +// for the typed memory view. +#define X(type, prefix, var, dim0, dim1) \ + auto var() const { \ + if constexpr (std::string_view(#dim0) == "nq") { \ + int start = model_->jnt_qposadr[id_]; \ + int end = (id_ < model_->njnt - 1) ? model_->jnt_qposadr[id_ + 1] : model_->nq; \ + return emscripten::val(emscripten::typed_memory_view(end - start, model_->prefix##var + start)); \ + } else if constexpr (std::string_view(#dim0) == "nv") { \ + int start = model_->jnt_dofadr[id_]; \ + int end = (id_ < model_->njnt - 1) ? model_->jnt_dofadr[id_ + 1] : model_->nv; \ + return emscripten::val(emscripten::typed_memory_view(end - start, model_->prefix##var + start)); \ + } else if constexpr (std::string_view(#dim0) == "nhfielddata") { \ + int start = model_->hfield_adr[id_]; \ + int count = model_->hfield_nrow[id_] * model_->hfield_ncol[id_]; \ + return emscripten::val(emscripten::typed_memory_view(count, model_->hfield_data + start)); \ + } else if constexpr (std::string_view(#dim0) == "ntexdata") { \ + int start = model_->tex_adr[id_]; \ + int count = model_->tex_height[id_] * model_->tex_width[id_] * model_->tex_nchannel[id_]; \ + return emscripten::val(emscripten::typed_memory_view(count, model_->tex_data + start)); \ + } else if constexpr (std::string_view(#dim0) == "nnumericdata") { \ + int start = model_->numeric_adr[id_]; \ + int count = model_->numeric_size[id_]; \ + return emscripten::val(emscripten::typed_memory_view(count, model_->numeric_data + start)); \ + } else if constexpr (std::string_view(#dim0) == "ntupledata") { \ + int start = model_->tuple_adr[id_]; \ + int count = model_->tuple_size[id_]; \ + return emscripten::val(emscripten::typed_memory_view(count, model_->prefix##var + start)); \ + } else { \ + if constexpr (std::string_view(#dim1) == "1") { \ + return model_->prefix##var[id_]; \ + } else { \ + return emscripten::val( \ + emscripten::typed_memory_view(dim1, model_->prefix##var + id_ * dim1)); \ + } \ + } \ + } + +// Expands to a struct definition for each object type in MJMODEL_ACCESSORS. +// Each struct, named `MjModel{Name}Accessor`, provides: +// - A constructor taking an `mjModel*` and an integer `id`. +// - An `id()` method to get the object's index. +// - A `name()` method to get the object's name using `mj_id2name`. +// - Member functions generated by the `MJMODEL_##NAME` macro, which in turn +// uses the `X` macro to define accessors for fields within `mjModel`. +#define X_ACCESSOR(NAME, Name, OBJTYPE, field, nfield) \ + struct MjModel##Name##Accessor { \ + MjModel##Name##Accessor(mjModel* model, int id) : model_(model), id_(id) {} \ + \ + int id() const { return id_; } \ + std::string name() const { \ + return mj_id2name(model_, OBJTYPE, id_); \ + } \ + \ + MJMODEL_##NAME \ + \ + private: \ + mjModel* model_; \ + int id_; \ + }; +MJMODEL_ACCESSORS +#undef X_ACCESSOR + +#undef X +#undef MJ_M + // {{ STRUCTS_HEADER }} // {{ STRUCTS_SOURCE }} @@ -473,6 +583,20 @@ int mj_setLengthRange_wrapper(const MjModel& m, const MjData& d, int index, cons EMSCRIPTEN_BINDINGS(mujoco_bindings) { // {{ ENUM_BINDINGS }} + // Bindings for the MjModel accessor classes. + #define X(type, prefix, var, dim0, dim1) .property(#var, &Accessor::var) + #define X_ACCESSOR(NAME, Name, OBJTYPE, field, nfield) \ + { \ + using Accessor = MjModel##Name##Accessor; \ + emscripten::class_("MjModel" #Name "Accessor") \ + .property("id", &Accessor::id) \ + .property("name", &Accessor::name) \ + MJMODEL_##NAME; \ + } + MJMODEL_ACCESSORS + #undef X + #undef X_ACCESSOR + // {{ STRUCTS_BINDINGS }} emscripten::class_("MjVFS") @@ -527,6 +651,9 @@ EMSCRIPTEN_BINDINGS(mujoco_bindings) { emscripten::register_vector("MjvGeomVec"); // register_type() improves type information (val is mapped to any by default) + // NumberOrString is used in functions returning accessors, allowing users to + // get an accessor by name (string) or id (number). + emscripten::register_type("number|string"); emscripten::register_type("number[]"); emscripten::register_type("string"); diff --git a/wasm/tests/bindings_test.ts b/wasm/tests/bindings_test.ts index f0c8a7a3..5501fbe2 100644 --- a/wasm/tests/bindings_test.ts +++ b/wasm/tests/bindings_test.ts @@ -75,14 +75,14 @@ function norm(arr: number[]): number { return Math.sqrt(arr.reduce((acc, val) => acc + val * val, 0)); } -function expectArraysClose(arr1: TypedArray, arr2: TypedArray, precision = 1) { +function expectArraysClose(arr1: any, arr2: TypedArray, precision = 1) { expect(arr1.length).toEqual(arr2.length); for (let i = 0; i < arr1.length; i++) { expect(arr1[i]).toBeCloseTo(arr2[i], precision); } } -function expectArraysEqual(arr1: TypedArray, arr2: TypedArray) { +function expectArraysEqual(arr1: any, arr2: TypedArray) { expect(arr1.length).toEqual(arr2.length); for (let i = 0; i < arr1.length; i++) { expect(arr1[i]).toEqual(arr2[i]); @@ -1872,4 +1872,272 @@ describe('MuJoCo WASM Bindings', () => { model?.delete(); } }); + + describe('MjModel named access', () => { + // Corresponds to + // bindings_test.py:test_named_indexing_invalid_names_in_model + it('should throw an error for invalid geom names in model', () => { + expect(() => { + model!.geom('badgeom'); + }).toThrowError('MuJoCo Error: Invalid name, MjModel.geom not found'); + }); + + // Corresponds to bindings_test.py:test_named_indexing_geom_size + it('should correctly access geom size using named indexing', () => { + const boxId = + mujoco.mj_name2id(model!, mujoco.mjtObj.mjOBJ_GEOM.value, 'mybox'); + assertExists(boxId); + + // Check that named indexing returns the same object as id indexing. + expect(model!.geom('mybox')).toEqual(model!.geom(boxId)); + expect(model!.geom('mybox')!.size).toEqual(model!.geom(boxId)!.size); + expect(model!.geom('mybox')!.size.length).toEqual(3); + + // Test that the indexer is returning a view into the underlying struct. + const sizeFromIndexer = model!.geom('mybox')!.size; + const originalGeomSize = new Float64Array(model!.geom_size); + + model!.geom_size.set([7, 11, 13], boxId * 3); + expectArraysEqual(sizeFromIndexer, new Float64Array([7, 11, 13])); + + model!.geom('mybox')!.size.set([5, 3, 2]); + expectArraysEqual( + model!.geom_size.slice(boxId * 3, boxId * 3 + 3), + new Float64Array([5, 3, 2])); + }); + + // Corresponds to bindings_test.py:test_named_indexing_geom_quat + it('should correctly access geom quat using named indexing', () => { + const boxId = + mujoco.mj_name2id(model!, mujoco.mjtObj.mjOBJ_GEOM.value, 'mybox'); + assertExists(boxId); + + // Check that named indexing returns the same object as id indexing. + expect(model!.geom('mybox')).toEqual(model!.geom(boxId)); + expect(model!.geom('mybox')!.quat).toEqual(model!.geom(boxId)!.quat); + expect(model!.geom('mybox')!.quat.length).toEqual(4); + + // Test that the indexer is returning a view into the underlying struct. + const quatFromIndexer = model!.geom('mybox')!.quat; + + model!.geom_quat.set([5, 10, 15, 20], boxId * 4); + expectArraysEqual(quatFromIndexer, new Float64Array([5, 10, 15, 20])); + + model!.geom('mybox')!.quat.set([12, 9, 6, 3]); + expectArraysEqual( + model!.geom_quat.slice(boxId * 4, boxId * 4 + 4), + new Float64Array([12, 9, 6, 3])); + }); + + it('should support named access for Joints', () => { + const xml = ` + `; + const tempXmlFilename = '/tmp/jnt.xml'; + writeXMLFile(tempXmlFilename, xml); + const model = mujoco.MjModel.mj_loadXML(tempXmlFilename); + try { + assertExists(model); + const j0 = model.jnt('root'); + expect(j0.id).toBe(0); + expect(j0.name).toBe('root'); + expect(j0.type).toBe(mujoco.mjtJoint.mjJNT_FREE.value); + expect(j0.qposadr).toBe(0); + expect(j0.dofadr).toBe(0); + expect(j0.axis).toEqual(new Float64Array([0, 0, 1])); + expect(j0.type).toEqual(mujoco.mjtJoint.mjJNT_FREE.value); + + const j0qposStart = model.jnt_qposadr[j0.id]; + const j0qposEnd = + (j0.id < model.njnt - 1) ? model.jnt_qposadr[j0.id + 1] : model.nq; + const expectedJ0qpos = model.qpos0.slice(j0qposStart, j0qposEnd); + expectArraysEqual( + expectedJ0qpos, new Float64Array([0, 0, 1, 1, 0, 0, 0])); + expectArraysEqual(j0.qpos0, expectedJ0qpos); + + const j0dofStart = model.jnt_dofadr[j0.id]; + const j0dofEnd = + (j0.id < model.njnt - 1) ? model.jnt_dofadr[j0.id + 1] : model.nv; + const expectedJ0dof = model.dof_bodyid.slice(j0dofStart, j0dofEnd); + expectArraysEqual(expectedJ0dof, new Int32Array([1, 1, 1, 1, 1, 1])); + expectArraysEqual(j0.bodyid, expectedJ0dof); + + + const j1 = model.jnt('hinge'); + expect(j1.id).toBe(1); + expect(j1.name).toBe('hinge'); + expect(j1.type).toBe(mujoco.mjtJoint.mjJNT_HINGE.value); + expect(j1.qposadr).toBe(7); + expect(j1.dofadr).toBe(6); + expect(j1.axis).toEqual(new Float64Array([0, 1, 0])); + expect(j1.frictionloss).toEqual(new Float64Array([0.4])); + + const j1qposStart = model.jnt_qposadr[j1.id]; + const j1qposEnd = + (j1.id < model.njnt - 1) ? model.jnt_qposadr[j1.id + 1] : model.nq; + const expectedJ1qpos = model.qpos0.slice(j1qposStart, j1qposEnd); + expectArraysEqual(expectedJ1qpos, new Float64Array([0])); + expectArraysEqual(j1.qpos0, expectedJ1qpos); + + const j1dofStart = model.jnt_dofadr[j1.id]; + const j1dofEnd = + (j1.id < model.njnt - 1) ? model.jnt_dofadr[j1.id + 1] : model.nv; + const expectedJ1dof = model.dof_bodyid.slice(j1dofStart, j1dofEnd); + expectArraysEqual(expectedJ1dof, new Int32Array([2])); + expectArraysEqual(j1.bodyid, new Int32Array([2])); + } finally { + model?.delete(); + unlinkXMLFile(tempXmlFilename); + } + }); + + it('should support named access for HFields', () => { + const xml = ` + + + + + + + `; + const tempXmlFilename = '/tmp/hf.xml'; + writeXMLFile(tempXmlFilename, xml); + const model = mujoco.MjModel.mj_loadXML(tempXmlFilename); + try { + assertExists(model); + const hf = model.hfield('hf'); + expect(hf.id).toBe(0); + expect(hf.name).toBe('hf'); + expect(hf.nrow).toBe(2); + expect(hf.ncol).toBe(3); + expectArraysClose(hf.size, new Float64Array([1, 2, 0.1, 0.2])); + const expectedDataLength = + model.hfield_nrow[hf.id] * model.hfield_ncol[hf.id]; + expect(expectedDataLength).toBe(6); + expect(hf.data.length).toBe(expectedDataLength); + } finally { + model?.delete(); + unlinkXMLFile(tempXmlFilename); + } + }); + + it('should support named access for Numeric', () => { + const xml = ` + + + + + `; + const tempXmlFilename = '/tmp/num.xml'; + writeXMLFile(tempXmlFilename, xml); + const model = mujoco.MjModel.mj_loadXML(tempXmlFilename); + try { + assertExists(model); + const n1 = model.numeric('n1'); + const expectedN1data = model.numeric_data.slice( + model.numeric_adr[n1.id], + model.numeric_adr[n1.id] + model.numeric_size[n1.id]); + expect(n1.id).toBe(0); + expect(n1.name).toBe('n1'); + expect(n1.size).toBe(3); + expect(expectedN1data).toEqual(new Float64Array([1, 2, 3])); + expectArraysClose(n1.data, expectedN1data); + + const n2 = model.numeric('n2'); + const expectedN2data = model.numeric_data.slice( + model.numeric_adr[n2.id], + model.numeric_adr[n2.id] + model.numeric_size[n2.id]); + expect(n2.id).toBe(1); + expect(n2.name).toBe('n2'); + expect(n2.size).toBe(4); + expect(expectedN2data).toEqual(new Float64Array([4, 5, 6, 7])); + expectArraysClose(n2.data, expectedN2data); + } finally { + model?.delete(); + unlinkXMLFile(tempXmlFilename); + } + }); + + it('should support named access for Texture', () => { + const xml = ` + + + + `; + const tempXmlFilename = '/tmp/tex.xml'; + writeXMLFile(tempXmlFilename, xml); + const model = mujoco.MjModel.mj_loadXML(tempXmlFilename); + try { + assertExists(model); + const t1 = model.tex('t1'); + const expectedData = model.tex_data.slice( + model.tex_adr[t1.id], + model.tex_adr[t1.id] + + model.tex_height[t1.id] * model.tex_width[t1.id] * + model.tex_nchannel[t1.id]); + + expect(t1.id).toBe(0); + expect(t1.name).toBe('t1'); + expect(t1.type).toBe(mujoco.mjtTexture.mjTEXTURE_2D.value); + expect(t1.width).toBe(10); + expect(t1.height).toBe(12); + expect(t1.nchannel).toBe(3); + expect(t1.data.length).toBe(t1.width * t1.height * t1.nchannel); + expectArraysClose(t1.data, expectedData); + } finally { + model?.delete(); + unlinkXMLFile(tempXmlFilename); + } + }); + + it('should support named access for Tuple', () => { + const xml = ` + + + + + + + + + + + `; + const tempXmlFilename = '/tmp/tup.xml'; + writeXMLFile(tempXmlFilename, xml); + const model = mujoco.MjModel.mj_loadXML(tempXmlFilename); + try { + assertExists(model); + const t1 = model.tuple('tup1'); + const expectedObjtype = model.tuple_objtype.slice( + model.tuple_adr[t1.id], + model.tuple_adr[t1.id] + model.tuple_size[t1.id]); + expect(t1.id).toBe(0); + expect(t1.name).toBe('tup1'); + expect(t1.size).toBe(2); + expectArraysEqual( + expectedObjtype, new Int32Array([ + mujoco.mjtObj.mjOBJ_GEOM.value, mujoco.mjtObj.mjOBJ_GEOM.value + ])); + expectArraysEqual(t1.objtype, expectedObjtype); + } finally { + model?.delete(); + unlinkXMLFile(tempXmlFilename); + } + }); + }); });