From 7dfd92098fc3248e144f71347a2a05bccc2c0e9d Mon Sep 17 00:00:00 2001 From: Matija Kecman Date: Thu, 6 Nov 2025 10:09:06 -0800 Subject: [PATCH] Simplified struct field handling in the WASM bindings PiperOrigin-RevId: 829003969 Change-Id: I139f9c0dad96e1664b1bd054f09fcdf58aff260e --- wasm/codegen/generated/bindings.cc | 63 ++-- wasm/codegen/helpers/common.py | 9 +- wasm/codegen/helpers/constants.py | 67 ++-- wasm/codegen/helpers/helpers_test.py | 44 +-- wasm/codegen/helpers/structs.py | 492 ++++++++++----------------- 5 files changed, 266 insertions(+), 409 deletions(-) diff --git a/wasm/codegen/generated/bindings.cc b/wasm/codegen/generated/bindings.cc index aa64d587..15e24fcc 100644 --- a/wasm/codegen/generated/bindings.cc +++ b/wasm/codegen/generated/bindings.cc @@ -4349,6 +4349,9 @@ struct MjModel { void set_nbuffer(mjtSize value) { ptr_->nbuffer = value; } + // opt field is handled manually in template file struct declaration + // vis field is handled manually in template file struct declaration + // stat field is handled manually in template file struct declaration emscripten::val buffer() const { return emscripten::val(emscripten::typed_memory_view(ptr_->nbuffer, static_cast(ptr_->buffer))); } @@ -5729,7 +5732,7 @@ struct MjData { void set_maxuse_efc(int value) { ptr_->maxuse_efc = value; } - // array field is defined manually. solver + // solver field is handled manually in template file struct declaration emscripten::val solver_niter() const { return emscripten::val(emscripten::typed_memory_view(20, ptr_->solver_niter)); } @@ -5739,8 +5742,8 @@ struct MjData { emscripten::val solver_fwdinv() const { return emscripten::val(emscripten::typed_memory_view(2, ptr_->solver_fwdinv)); } - // array field is defined manually. warning - // array field is defined manually. timer + // warning field is handled manually in template file struct declaration + // timer field is handled manually in template file struct declaration int ncon() const { return ptr_->ncon; } @@ -6077,7 +6080,7 @@ struct MjData { emscripten::val cfrc_ext() const { return emscripten::val(emscripten::typed_memory_view(model->nbody * 6, ptr_->cfrc_ext)); } - // complex pointer field is defined manually. contact + // contact field is handled manually in template file struct declaration emscripten::val efc_type() const { return emscripten::val(emscripten::typed_memory_view(ptr_->nefc, ptr_->efc_type)); } @@ -6380,26 +6383,26 @@ struct MjvScene { void set_ngeom(int value) { ptr_->ngeom = value; } - // complex pointer field is defined manually. geoms - // primitive pointer field with complex extents is defined manually. geomorder + // geoms field is handled manually in template file struct declaration + // geomorder field is handled manually in template file struct declaration int nflex() const { return ptr_->nflex; } void set_nflex(int value) { ptr_->nflex = value; } - // primitive pointer field with complex extents is defined manually. flexedgeadr - // primitive pointer field with complex extents is defined manually. flexedgenum - // primitive pointer field with complex extents is defined manually. flexvertadr - // primitive pointer field with complex extents is defined manually. flexvertnum - // primitive pointer field with complex extents is defined manually. flexfaceadr - // primitive pointer field with complex extents is defined manually. flexfacenum - // primitive pointer field with complex extents is defined manually. flexfaceused - // primitive pointer field with complex extents is defined manually. flexedge - // primitive pointer field with complex extents is defined manually. flexvert - // primitive pointer field with complex extents is defined manually. flexface - // primitive pointer field with complex extents is defined manually. flexnormal - // primitive pointer field with complex extents is defined manually. flextexcoord + // flexedgeadr field is handled manually in template file struct declaration + // flexedgenum field is handled manually in template file struct declaration + // flexvertadr field is handled manually in template file struct declaration + // flexvertnum field is handled manually in template file struct declaration + // flexfaceadr field is handled manually in template file struct declaration + // flexfacenum field is handled manually in template file struct declaration + // flexfaceused field is handled manually in template file struct declaration + // flexedge field is handled manually in template file struct declaration + // flexvert field is handled manually in template file struct declaration + // flexface field is handled manually in template file struct declaration + // flexnormal field is handled manually in template file struct declaration + // flextexcoord field is handled manually in template file struct declaration mjtByte flexvertopt() const { return ptr_->flexvertopt; } @@ -6430,19 +6433,19 @@ struct MjvScene { void set_nskin(int value) { ptr_->nskin = value; } - // primitive pointer field with complex extents is defined manually. skinfacenum - // primitive pointer field with complex extents is defined manually. skinvertadr - // primitive pointer field with complex extents is defined manually. skinvertnum - // primitive pointer field with complex extents is defined manually. skinvert - // primitive pointer field with complex extents is defined manually. skinnormal + // skinfacenum field is handled manually in template file struct declaration + // skinvertadr field is handled manually in template file struct declaration + // skinvertnum field is handled manually in template file struct declaration + // skinvert field is handled manually in template file struct declaration + // skinnormal field is handled manually in template file struct declaration int nlight() const { return ptr_->nlight; } void set_nlight(int value) { ptr_->nlight = value; } - // array field is defined manually. lights - // array field is defined manually. camera + // lights field is handled manually in template file struct declaration + // camera field is handled manually in template file struct declaration mjtByte enabletransform() const { return ptr_->enabletransform; } @@ -6505,7 +6508,7 @@ struct MjSpec { std::unique_ptr copy(); mjSpec* get() const; void set(mjSpec* ptr); - // complex pointer field is defined manually. element + // element field is handled manually in template file struct declaration mjString modelname() const { return (ptr_ && ptr_->modelname) ? *(ptr_->modelname) : ""; } @@ -6514,16 +6517,16 @@ struct MjSpec { *(ptr_->modelname) = value; } } - // struct field is defined manually. compiler + // compiler field is handled manually in template file struct declaration mjtByte strippath() const { return ptr_->strippath; } void set_strippath(mjtByte value) { ptr_->strippath = value; } - // struct field is defined manually. option - // struct field is defined manually. visual - // struct field is defined manually. stat + // option field is handled manually in template file struct declaration + // visual field is handled manually in template file struct declaration + // stat field is handled manually in template file struct declaration mjtSize memory() const { return ptr_->memory; } diff --git a/wasm/codegen/helpers/common.py b/wasm/codegen/helpers/common.py index 84fc03c1..c809b6ac 100644 --- a/wasm/codegen/helpers/common.py +++ b/wasm/codegen/helpers/common.py @@ -61,11 +61,14 @@ def write_to_file(filepath: str, content: str) -> None: print(f"Error writing to output file: {filepath} - {e}") +def lowercase_first_letter(input_string: str) -> str: + """Lowercases the first letter of a string.""" + return input_string[:1].lower() + input_string[1:] + + def uppercase_first_letter(input_string: str) -> str: """Uppercases the first letter of a string.""" - if input_string: - return input_string[0].upper() + input_string[1:] - return input_string + return input_string[:1].upper() + input_string[1:] def try_cast_to_scalar_type(value: str) -> int | float | str: diff --git a/wasm/codegen/helpers/constants.py b/wasm/codegen/helpers/constants.py index 52007fbc..e24fcecd 100644 --- a/wasm/codegen/helpers/constants.py +++ b/wasm/codegen/helpers/constants.py @@ -307,6 +307,17 @@ SKIPPED_STRUCTS: List[str] = [ # go/keep-sorted end ] +# These structs require specific function calls for creation and/or deletion, +# or some of their fields need to be handled manually for now; +# making their wrapper constructors/destructors non-trivial. +MANUAL_STRUCTS: List[str] = [ + "MjData", + "MjModel", + "MjvScene", + "MjSpec", + "MjVisual", +] + # Dictionary that maps anonymous structs to their parent struct and field name. # Anonymous structs are not defined as independent structs in the MuJoCo # codebase, but they are part of other structs. This dictionary is used to @@ -325,9 +336,9 @@ ANONYMOUS_STRUCTS: Dict[str, Dict[str, str]] = { # This list is created by subtracting the skipped structs from the list of all # structs and adding the anonymous structs. STRUCTS_TO_BIND: List[str] = list( - (set(introspect_structs.STRUCTS.keys()) - set(SKIPPED_STRUCTS)).union( - ANONYMOUS_STRUCTS.keys() - ) + set(introspect_structs.STRUCTS.keys()) + .union(ANONYMOUS_STRUCTS.keys()) + .difference(set(SKIPPED_STRUCTS)) ) # List of structs that do not have a default constructor. @@ -421,37 +432,42 @@ MJDATA_SIZES: List[str] = [ # go/keep-sorted end ] -# Dictionary where keys are the struct names and the values are lists of the -# fields that are manually specified in the bindings.cc template file. -MANUALLY_ADDED_FIELDS_FROM_TEMPLATE: Dict[str, List[str]] = { +# Fields that should be entirely omitted from the bindings. +SKIPPED_FIELDS: Dict[str, List[str]] = {} + +# Fields handled manually in template file struct declaration. +MANUAL_FIELDS: Dict[str, List[str]] = { # go/keep-sorted start "MjData": ["solver", "timer", "warning", "contact"], + "MjModel": ["opt", "vis", "stat"], "MjSpec": ["option", "visual", "stat", "element", "compiler"], "MjvScene": [ - "model", - "lights", + # go/keep-sorted start "camera", - "geoms", - "geomorder", + "flexedge", "flexedgeadr", "flexedgenum", - "flexvertadr", - "flexvertnum", + "flexface", "flexfaceadr", "flexfacenum", "flexfaceused", - "flexedge", - "flexvert", - "skinfacenum", - "skinvertadr", - "skinvertnum", - "skinvert", - "skinnormal", - "flexface", "flexnormal", "flextexcoord", + "flexvert", + "flexvertadr", + "flexvertnum", + "geomorder", + "geoms", + "lights", + "model", + "skinfacenum", + "skinnormal", + "skinvert", + "skinvertadr", + "skinvertnum", # go/keep-sorted end ], + # go/keep-sorted end } # Dictionary that maps byte array fields to their corresponding size members. @@ -462,14 +478,3 @@ BYTE_FIELDS: Dict[str, Dict[str, str]] = { "buffer": {"size": "nbuffer"}, "arena": {"size": "narena"}, } - -# These structs require specific function calls for creation and/or deletion, -# or some of their fields need to be handled manually for now; -# making their wrapper constructors/destructors non-trivial. -HARDCODED_WRAPPER_STRUCTS: List[str] = [ - "MjData", - "MjModel", - "MjvScene", - "MjSpec", - "MjVisual", -] diff --git a/wasm/codegen/helpers/helpers_test.py b/wasm/codegen/helpers/helpers_test.py index 22bf222c..57f63f2f 100644 --- a/wasm/codegen/helpers/helpers_test.py +++ b/wasm/codegen/helpers/helpers_test.py @@ -349,9 +349,9 @@ void MjsElement::set(mjsElement* ptr) { ), doc="", ) - wrapped_field_data = structs.StructFieldHandler( + wrapped_field_data = structs._generate_field_data( field_with_init, "MjsTexture" - ).generate() + ) self.assertEqual( structs.build_struct_source( "mjsTexture", @@ -573,8 +573,7 @@ class StructFieldCodeBuilderTest(absltest.TestCase): type=ast_nodes.ValueType(name="int"), doc="number of geoms", ) - self.assertEqual( - structs.build_primitive_type_definition(field), + self.assertEqual(structs._generate_field_data(field, "ngeom").definition, """ int ngeom() const { return ptr_->ngeom; @@ -595,9 +594,7 @@ void set_ngeom(int value) { array_extent=("ngeom", 4), ) self.assertEqual( - structs.build_memory_view_definition( - field, "ptr_->ngeom * 4", "ptr_->geom_rgba" - ), + structs._generate_field_data(field, "geom_rgba").definition, """ emscripten::val geom_rgba() const { return emscripten::val(emscripten::typed_memory_view(ptr_->ngeom * 4, ptr_->geom_rgba)); @@ -614,7 +611,7 @@ emscripten::val geom_rgba() const { doc="rgba when material is omitted", ) self.assertEqual( - structs.build_string_field_definition(field), + structs._generate_field_data(field, "MjString").definition, """ mjString string_field() const { return (ptr_ && ptr_->string_field) ? *(ptr_->string_field) : ""; @@ -636,7 +633,7 @@ void set_string_field(const mjString& value) { doc="", ) self.assertEqual( - structs.build_mjvec_pointer_definition(field, "mjDoubleVec"), + structs._generate_field_data(field, "mjDoubleVec").definition, """ mjDoubleVec &vector_field() const { return *(ptr_->vector_field); @@ -652,7 +649,7 @@ mjDoubleVec &vector_field() const { doc="", ) self.assertEqual( - structs.build_mjvec_pointer_definition(field, "mjByteVec"), + structs._generate_field_data(field, "mjByteVec").definition, """ std::vector &vector_field() const { return *(reinterpret_cast*>(ptr_->vector_field)); @@ -666,7 +663,7 @@ std::vector &vector_field() const { doc="number of geoms", ) self.assertEqual( - structs.build_simple_property_binding(field, "MjModel"), + structs._simple_property_binding(field, "MjModel"), '.property("ngeom", &MjModel::ngeom)', ) @@ -677,7 +674,7 @@ std::vector &vector_field() const { doc="", ) self.assertEqual( - structs.build_simple_property_binding(field, "MjModel", True), + structs._simple_property_binding(field, "MjModel", True), '.property("ngeom", &MjModel::ngeom, &MjModel::set_ngeom)', ) @@ -688,11 +685,11 @@ std::vector &vector_field() const { doc="", ) self.assertEqual( - structs.build_simple_property_binding( + structs._simple_property_binding( field, "MjModel", - add_setter=True, - add_return_value_policy_as_ref=True, + setter=True, + reference=True, ), '.property("ngeom", &MjModel::ngeom, &MjModel::set_ngeom, reference())', ) @@ -708,8 +705,7 @@ class StructFieldHandlerTest(absltest.TestCase): doc="number of geoms", ) - field_handler_scalar = structs.StructFieldHandler(field_scalar, "MjModel") - wrapped_field_data = field_handler_scalar.generate() + wrapped_field_data = structs._generate_field_data(field_scalar, "MjModel") self.assertEqual( wrapped_field_data.definition, """ @@ -736,7 +732,7 @@ void set_ngeom(int value) { doc="rgba when material is omitted", array_extent=("ngeom", 4), ) - wrapped_field_data = structs.StructFieldHandler(field, "MjModel").generate() + wrapped_field_data = structs._generate_field_data(field, "MjModel") self.assertEqual( wrapped_field_data.definition, @@ -761,7 +757,7 @@ emscripten::val geom_rgba() const { ), doc="main buffer; all pointers point in it (nbuffer bytes)", ) - wrapped_field_data = structs.StructFieldHandler(field, "MjData").generate() + wrapped_field_data = structs._generate_field_data(field, "MjData") self.assertEqual( wrapped_field_data.definition, @@ -782,9 +778,7 @@ emscripten::val buffer() const { ), doc="", ) - wrapped_field_data = structs.StructFieldHandler( - field, "MjsTexture" - ).generate() + wrapped_field_data = structs._generate_field_data(field, "MjsTexture") self.assertEqual(wrapped_field_data.definition, "MjsElement element;") self.assertEqual( @@ -806,9 +800,7 @@ emscripten::val buffer() const { ), doc="gravitational acceleration", ) - wrapped_field_data = structs.StructFieldHandler( - field, "MjOption" - ).generate() + wrapped_field_data = structs._generate_field_data(field, "MjOption") self.assertEqual( wrapped_field_data.definition, @@ -833,7 +825,7 @@ emscripten::val gravity() const { ), doc="description", ) - wrapped_field_data = structs.StructFieldHandler(field, "MjModel").generate() + wrapped_field_data = structs._generate_field_data(field, "MjModel") self.assertEqual( wrapped_field_data.definition, """ diff --git a/wasm/codegen/helpers/structs.py b/wasm/codegen/helpers/structs.py index ec09573d..74981d17 100644 --- a/wasm/codegen/helpers/structs.py +++ b/wasm/codegen/helpers/structs.py @@ -49,7 +49,8 @@ class WrappedFieldData: # Whether the field is a primitive or fixed size is_primitive_or_fixed_size: bool = False - # Underlying type of the field + # Underlying type of the field. If non-empty, used to determine the order in + # which structs are written in the bindings.h file. typename: str = "" @@ -76,360 +77,217 @@ class WrappedStructData: use_shallow_copy: bool = True -def build_primitive_type_definition(field: ast_nodes.StructFieldDecl) -> str: - """Builds the C++ getter/setter code for a primitive type field wrapper.""" - if not isinstance(field.type, ast_nodes.ValueType): - raise ValueError(f"{field.type} must be ValueType.") - - builder = code_builder.CodeBuilder() - with builder.function(f"{field.type.name} {field.name}() const"): - builder.line(f"return ptr_->{field.name};") - with builder.function(f"void set_{field.name}({field.type.name} value)"): - builder.line(f"ptr_->{field.name} = value;") - return builder.to_string() - - -def build_memory_view_definition( - field: ast_nodes.StructFieldDecl, array_size_str: str, ptr_expr: str -) -> str: - """Builds the C++ code for a pointer type field wrapper.""" - builder = code_builder.CodeBuilder() - with builder.function(f"emscripten::val {field.name}() const"): - builder.line( - "return" - f" emscripten::val(emscripten::typed_memory_view({array_size_str}," - f" {ptr_expr}));" - ) - return builder.to_string() - - -def build_string_field_definition(field: ast_nodes.StructFieldDecl) -> str: - """Builds the C++ code getter/setter for a string type field wrapper.""" - builder = code_builder.CodeBuilder() - with builder.function(f"mjString {field.name}() const"): - builder.line( - f'return (ptr_ && ptr_->{field.name}) ? *(ptr_->{field.name}) : "";' - ) - with builder.function(f"void set_{field.name}(const mjString& value)"): - with builder.block(f"if (ptr_ && ptr_->{field.name})"): - builder.line(f"*(ptr_->{field.name}) = value;") - return builder.to_string() - - -def build_mjvec_pointer_definition( - field: ast_nodes.StructFieldDecl, vector_type: str -) -> str: - """Builds the C++ code for a mjVec type field wrapper.""" - ptr_field_expr = f"*(ptr_->{field.name})" - if vector_type == "mjByteVec": - vector_type = "std::vector" - ptr_field_expr = ( - f"*(reinterpret_cast*>(ptr_->{field.name}))" - ) - builder = code_builder.CodeBuilder() - with builder.function(f"{vector_type} &{field.name}() const"): - builder.line(f"return {ptr_field_expr};") - return builder.to_string() - - -def build_simple_property_binding( +def _simple_property_binding( field: ast_nodes.StructFieldDecl, struct_wrapper_name: str, - add_setter: bool = False, - add_return_value_policy_as_ref: bool = False, + setter: bool = False, + reference: bool = False, ) -> str: """Builds the C++ code for a simple property binding.""" - builder = code_builder.CodeBuilder() - setter_txt = "" - if add_setter: - setter_txt = f", &{struct_wrapper_name}::set_{field.name}" - if add_return_value_policy_as_ref: - as_reference_txt = ", reference()" - else: - as_reference_txt = "" - builder.line( - f'.property("{field.name}",' - f" &{struct_wrapper_name}::{field.name}{setter_txt}{as_reference_txt})" - ) - return builder.to_string() + f = field + w = struct_wrapper_name + setter_txt = f", &{w}::set_{f.name}" if setter else "" + reference_txt = ", reference()" if reference else "" + return f'.property("{f.name}", &{w}::{f.name}{setter_txt}{reference_txt})' -class StructFieldHandler: - """Class to handle the different struct field types, and provide the c++ code for the definitions and bindings.""" +def _generate_field_data( + field: ast_nodes.StructFieldDecl, struct_wrapper_name: str +) -> WrappedFieldData: + """Generates the C++ definition and binding code for the struct field.""" + f = field + w = struct_wrapper_name + s = common.lowercase_first_letter(w) - def __init__( - self, - field: ast_nodes.StructFieldDecl, - struct_wrapper_name: str, - ): - self.field = field - self.struct_wrapper_name = struct_wrapper_name - self.simple_property_binding = build_simple_property_binding( - self.field, self.struct_wrapper_name - ) - self.manually_added_fields = ( - constants.MANUALLY_ADDED_FIELDS_FROM_TEMPLATE.get( - self.struct_wrapper_name, {} - ) - ) - - def generate(self) -> WrappedFieldData: - """Generates the C++ definition and binding code for the struct field.""" - field_type = self.field.type - if isinstance(field_type, ast_nodes.ValueType) and ( - field_type.name in constants.PRIMITIVE_TYPES - or field_type.name.startswith("mjt") - ): - return self._handle_primitive() - elif isinstance(field_type, ast_nodes.PointerType): - return self._handle_pointer() - elif isinstance(field_type, ast_nodes.ArrayType): - return self._handle_array() - elif isinstance( - field_type, ast_nodes.ValueType - ) and field_type.name.startswith("mj"): - return self._handle_mj_struct() - elif isinstance(field_type, ast_nodes.AnonymousStructDecl): - return self._handle_anonymous_struct() - return self._undefined() - - def _handle_primitive(self) -> WrappedFieldData: - """Handles the generation of C++ definition and binding code for primitive fields.""" + if f.name in constants.MANUAL_FIELDS.get(w, []): + # Note: Manually handled MjModel fields are special cased so that a + # by-reference embind return value policy is used. return WrappedFieldData( - definition=(build_primitive_type_definition(self.field)), - typename=_get_field_struct_type(self.field.type), - binding=build_simple_property_binding( - self.field, - self.struct_wrapper_name, - add_setter=True, - add_return_value_policy_as_ref=True, - ), + typename=_get_field_struct_type(f.type), + definition=f"// {f.name} field is handled manually in template file struct declaration", # pylint: disable=line-too-long + binding=_simple_property_binding(f, w, reference=(w == "MjModel")), + ) + + if f.name in constants.SKIPPED_FIELDS.get(w, []): + return WrappedFieldData( + typename="", + definition=f"// {f.name} field is skipped.", + binding=f"// {f.name} field is skipped.", + ) + + if isinstance(f.type, ast_nodes.ValueType) and ( + f.type.name in constants.PRIMITIVE_TYPES or f.type.name.startswith("mjt") + ): + + builder = code_builder.CodeBuilder() + with builder.function(f"{f.type.name} {f.name}() const"): + builder.line(f"return ptr_->{f.name};") + with builder.function(f"void set_{f.name}({f.type.name} value)"): + builder.line(f"ptr_->{f.name} = value;") + + return WrappedFieldData( + definition=builder.to_string(), + typename=_get_field_struct_type(f.type), + binding=_simple_property_binding(f, w, setter=True, reference=True), is_primitive_or_fixed_size=True, ) - def _handle_pointer(self) -> WrappedFieldData: - """Handles the generation of C++ definition and binding code for pointer fields.""" - if not isinstance(self.field.type, ast_nodes.PointerType): - raise ValueError( - f"Expected PointerType, got {type(self.field.type)} for field" - f" {self.field.name}" + elif isinstance(f.type, ast_nodes.ValueType) and f.type.name.startswith("mj"): + + return WrappedFieldData( + definition=f"{common.uppercase_first_letter(f.type.name)} {f.name};", + typename=_get_field_struct_type(f.type), + binding=_simple_property_binding(f, w, setter=False, reference=True), + initialization=f", {f.name}(&ptr_->{f.name})", + ptr_copy_reset=f"{f.name}.set(&ptr_->{f.name});", + is_primitive_or_fixed_size=True, + ) + + elif isinstance(f.type, ast_nodes.AnonymousStructDecl): + + anonymous_struct_name = "" + for name, value in constants.ANONYMOUS_STRUCTS.items(): + if (value["parent"] == s and value["field_name"] == f.name): + anonymous_struct_name = name + break + + if anonymous_struct_name in constants.STRUCTS_TO_BIND: + return WrappedFieldData( + binding=_simple_property_binding(f, w, setter=False, reference=True), + typename=_get_field_struct_type(f.type), + initialization=f", {f.name}(&ptr_->{f.name})", + ptr_copy_reset=f"{f.name}.set(&ptr_->{f.name});", + is_primitive_or_fixed_size=True, ) - field_type: ast_nodes.PointerType = self.field.type + + elif isinstance(f.type, ast_nodes.ArrayType): + + inner_type = f.type.inner_type + size = math.prod(f.type.extents) + + if ( + isinstance(inner_type, ast_nodes.ValueType) + and inner_type.name in constants.PRIMITIVE_TYPES + ): + ptr_expr = f"ptr_->{f.name}" + if len(f.type.extents) > 1: + # for multi-dimensional arrays, we need to cast the field + # to a pointer, so embind can correctly interpret the memory + # view + ptr_expr = f"reinterpret_cast<{inner_type.name}*>({ptr_expr})" + + builder = code_builder.CodeBuilder() + with builder.function(f"emscripten::val {f.name}() const"): + builder.line( + "return" + f" emscripten::val(emscripten::typed_memory_view({str(size)}," + f" {ptr_expr}));" + ) + + return WrappedFieldData( + definition=builder.to_string(), + typename=_get_field_struct_type(f.type), + binding=_simple_property_binding(f, w), + is_primitive_or_fixed_size=True, + ) + else: + return WrappedFieldData( + definition=f"// TODO: NOT IMPLEMENTED ARRAY wrapper for {f.name}", + typename=_get_field_struct_type(f.type), + binding=f"// TODO: NOT IMPLEMENTED ARRAY binding for {f.name}", + ) + + elif isinstance(f.type, ast_nodes.PointerType): + inner_type_name = ( - field_type.inner_type.name - if isinstance(field_type.inner_type, ast_nodes.ValueType) + f.type.inner_type.name + if isinstance(f.type.inner_type, ast_nodes.ValueType) else "" ) - ptr_field_expr = f"ptr_->{self.field.name}" + ptr_field_expr = f"ptr_->{f.name}" array_size_str = "" - if self.field.array_extent: - array_size_str = parse_array_extent( - self.field.array_extent, self.struct_wrapper_name, self.field.name - ) - elif self.field.name in constants.BYTE_FIELDS.keys(): + if f.array_extent: + array_size_str = parse_array_extent(f.array_extent, w, f.name) + elif f.name in constants.BYTE_FIELDS.keys(): # for byte fields, we need to cast the pointer to uint8_t* # so embind can correctly interpret the memory view ptr_field_expr = f"static_cast({ptr_field_expr})" # for these byte fields, there is no array_extent, so we add the size of # in the config file based in the documentation - extent = (constants.BYTE_FIELDS[self.field.name]["size"],) - array_size_str = parse_array_extent( - extent, self.struct_wrapper_name, self.field.name - ) + extent = (constants.BYTE_FIELDS[f.name]["size"],) + array_size_str = parse_array_extent(extent, w, f.name) elif inner_type_name == "mjString": + + builder = code_builder.CodeBuilder() + with builder.function(f"mjString {f.name}() const"): + builder.line( + f'return (ptr_ && ptr_->{f.name}) ? *(ptr_->{f.name}) : "";' + ) + with builder.function(f"void set_{f.name}(const mjString& value)"): + with builder.block(f"if (ptr_ && ptr_->{f.name})"): + builder.line(f"*(ptr_->{f.name}) = value;") + return WrappedFieldData( - definition=build_string_field_definition(self.field), - typename=_get_field_struct_type(self.field.type), - binding=build_simple_property_binding( - self.field, - self.struct_wrapper_name, - add_setter=True, - add_return_value_policy_as_ref=True, - ), + definition=builder.to_string(), + typename=_get_field_struct_type(f.type), + binding=_simple_property_binding(f, w, setter=True, reference=True), ) elif inner_type_name.startswith("mj") and inner_type_name.endswith("Vec"): + ptr_field_expr_vec = f"*(ptr_->{f.name})" + vector_type = inner_type_name + if vector_type == "mjByteVec": + vector_type = "std::vector" + ptr_field_expr_vec = ( + f"*(reinterpret_cast*>(ptr_->{f.name}))" + ) + + builder = code_builder.CodeBuilder() + with builder.function(f"{vector_type} &{f.name}() const"): + builder.line(f"return {ptr_field_expr_vec};") + return WrappedFieldData( - definition=build_mjvec_pointer_definition( - self.field, inner_type_name - ), - typename=_get_field_struct_type(self.field.type), - binding=build_simple_property_binding( - self.field, - self.struct_wrapper_name, - add_setter=False, - add_return_value_policy_as_ref=True, - ), - ) - elif inner_type_name in constants.PRIMITIVE_TYPES: - return self._get_manual_definition( - comment_type="primitive pointer field with complex extents" + definition=builder.to_string(), + typename=_get_field_struct_type(f.type), + binding=_simple_property_binding(f, w, setter=False, reference=True) ) if ( inner_type_name.startswith("mj") and inner_type_name not in constants.PRIMITIVE_TYPES + and not f.array_extent + and w not in constants.MANUAL_FIELDS.keys() ): - # it's a pointer to a single struct, - # like the `element` field in mjs structs - # and the struct is not manually added - if ( - not self.field.array_extent - and self.struct_wrapper_name - not in constants.MANUALLY_ADDED_FIELDS_FROM_TEMPLATE.keys() - ): - ptr_field = cast(ast_nodes.PointerType, self.field.type) - wrapper_field_name = common.uppercase_first_letter( - cast(ast_nodes.ValueType, ptr_field.inner_type).name - ) - return WrappedFieldData( - definition=f"{wrapper_field_name} {self.field.name};", - typename=_get_field_struct_type(self.field.type), - binding=build_simple_property_binding( - self.field, - self.struct_wrapper_name, - add_setter=False, - add_return_value_policy_as_ref=True, - ), - initialization=f", {self.field.name}(ptr_->{self.field.name})", - ) - else: - return self._get_manual_definition(comment_type="complex pointer field") - - return WrappedFieldData( - typename=_get_field_struct_type(self.field.type), - definition=( - build_memory_view_definition( - self.field, array_size_str, ptr_field_expr - ) - ), - binding=self.simple_property_binding, - ) - - def _handle_array(self) -> WrappedFieldData: - """Handles the generation of C++ definition and binding code for array fields.""" - field_type = self.field.type - if not isinstance(field_type, ast_nodes.ArrayType): - raise ValueError( - f"Expected ArrayType, got {type(field_type)} for field" - f" {self.field.name}" + ptr_field = cast(ast_nodes.PointerType, f.type) + wrapper_field_name = common.uppercase_first_letter( + cast(ast_nodes.ValueType, ptr_field.inner_type).name ) - inner_type = field_type.inner_type - size = math.prod(field_type.extents) - - if isinstance(inner_type, ast_nodes.ValueType): - if inner_type.name in constants.PRIMITIVE_TYPES: - ptr_expr = f"ptr_->{self.field.name}" - if len(field_type.extents) > 1: - # for multi-dimensional arrays, we need to cast the field - # to a pointer, so embind can correctly interpret the memory - # view - ptr_expr = f"reinterpret_cast<{inner_type.name}*>({ptr_expr})" - return WrappedFieldData( - definition=( - build_memory_view_definition(self.field, str(size), ptr_expr) - ), - typename=_get_field_struct_type(self.field.type), - binding=self.simple_property_binding, - is_primitive_or_fixed_size=True, - ) - elif inner_type.name.startswith("mj") and not inner_type.name.startswith( - "mjt" - ): - return self._get_manual_definition(comment_type="array field") - - return WrappedFieldData( - definition=( - f"// TODO: NOT IMPLEMENTED ARRAY wrapper for {self.field.name}" - ), - typename=_get_field_struct_type(self.field.type), - binding=f"// TODO: NOT IMPLEMENTED ARRAY binding for {self.field.name}", - ) - - def _handle_mj_struct(self) -> WrappedFieldData: - """Handles the generation of C++ definition and binding code for mj struct fields.""" - if ( - isinstance(self.field.type, ast_nodes.ValueType) - and self.field.name not in self.manually_added_fields - and self.field.type.name in constants.STRUCTS_TO_BIND - ): - # TODO(manevi): Find a better way to do this instead of checking the - # struct wrapper name. - definition = "" - if self.struct_wrapper_name not in constants.HARDCODED_WRAPPER_STRUCTS: - wrapper_field_name = common.uppercase_first_letter(self.field.type.name) - definition = f"{wrapper_field_name} {self.field.name};" return WrappedFieldData( - definition=definition, - typename=_get_field_struct_type(self.field.type), - binding=build_simple_property_binding( - self.field, - self.struct_wrapper_name, - add_setter=False, - add_return_value_policy_as_ref=True, - ), - initialization=f", {self.field.name}(&ptr_->{self.field.name})", - ptr_copy_reset=f"{self.field.name}.set(&ptr_->{self.field.name});", - is_primitive_or_fixed_size=True, + definition=f"{wrapper_field_name} {f.name};", + typename=_get_field_struct_type(f.type), + binding=_simple_property_binding(f, w, setter=False, reference=True), + initialization=f", {f.name}(ptr_->{f.name})", ) - return self._get_manual_definition(comment_type="struct field") - def _handle_anonymous_struct(self) -> WrappedFieldData: - """Handles the generation of C++ definition and binding code for anonymous struct fields.""" - - anonymous_struct_name = "" - for name, value in constants.ANONYMOUS_STRUCTS.items(): - if ( - common.uppercase_first_letter(value["parent"]) - == self.struct_wrapper_name - and value["field_name"] == self.field.name - ): - anonymous_struct_name = name - break - - if ( - isinstance(self.field.type, ast_nodes.AnonymousStructDecl) - and self.field.name not in self.manually_added_fields - and anonymous_struct_name in constants.STRUCTS_TO_BIND - ): - return WrappedFieldData( - binding=build_simple_property_binding( - self.field, - self.struct_wrapper_name, - add_setter=False, - add_return_value_policy_as_ref=True, - ), - typename=_get_field_struct_type(self.field.type), - initialization=f", {self.field.name}(&ptr_->{self.field.name})", - ptr_copy_reset=f"{self.field.name}.set(&ptr_->{self.field.name});", - is_primitive_or_fixed_size=True, + builder = code_builder.CodeBuilder() + with builder.function(f"emscripten::val {f.name}() const"): + builder.line( + "return" + f" emscripten::val(emscripten::typed_memory_view({array_size_str}," + f" {ptr_field_expr}));" ) - return self._get_manual_definition(comment_type="anonymous struct field") - def _undefined(self) -> WrappedFieldData: - """This function adds a TODO comment for fields that are not handled by this class yet.""" return WrappedFieldData( - definition=f"// TODO: UNDEFINED definition for {self.field.name}", - binding=f"// TODO: UNDEFINED binding for {self.field.name}", + definition=builder.to_string(), + typename=_get_field_struct_type(f.type), + binding=_simple_property_binding(f, w), ) - def _get_manual_definition(self, comment_type: str = "") -> WrappedFieldData: - """Helper method to generate a comment as a definition for manually added fields.""" - if self.field.name in self.manually_added_fields: - return WrappedFieldData( - typename=_get_field_struct_type(self.field.type), - definition=( - f"// {comment_type} is defined manually. {self.field.name}" - ), - binding=self.simple_property_binding, - ) - return WrappedFieldData( - typename=_get_field_struct_type(self.field.type), - definition=( - f"// TODO: Define {comment_type} manually for {self.field.name}" - ), - binding=f"// TODO: {self.simple_property_binding}", - ) + return WrappedFieldData( + definition=f"// TODO: UNDEFINED definition for {f.name}", + typename=_get_field_struct_type(f.type), + binding=f"// TODO: UNDEFINED binding for {f.name}", + ) def _has_nested_wrapper_members(struct_info: ast_nodes.StructDecl) -> bool: @@ -502,10 +360,6 @@ def _default_function_statement(struct_name: str) -> str: if ( struct_name in constants.ANONYMOUS_STRUCTS.keys() or struct_name in constants.NO_DEFAULT_CONSTRUCTORS - or ( - common.uppercase_first_letter(struct_name) - in constants.MANUALLY_ADDED_FIELDS_FROM_TEMPLATE.keys() - ) ): return "" elif struct_name.startswith("mjs"): @@ -568,7 +422,7 @@ def build_struct_header( is_anonymous_struct = struct_name in constants.ANONYMOUS_STRUCTS.keys() is_hardcoded_wrapper_struct = ( common.uppercase_first_letter(struct_name) - in constants.HARDCODED_WRAPPER_STRUCTS + in constants.MANUAL_STRUCTS ) if ( @@ -765,7 +619,7 @@ def generate_wasm_bindings( wrapped_fields: List[WrappedFieldData] = [] for field in struct_fields: - wrapped_field = StructFieldHandler(field, wrapped_name).generate() + wrapped_field = _generate_field_data(field, wrapped_name) wrapped_fields.append(wrapped_field) wrapped_header = build_struct_header(