diff --git a/wasm/codegen/generators/structs.py b/wasm/codegen/generators/structs.py index e5baa770..bea5b060 100644 --- a/wasm/codegen/generators/structs.py +++ b/wasm/codegen/generators/structs.py @@ -17,7 +17,7 @@ from typing import Optional from wasm.codegen.helpers import constants -from wasm.codegen.helpers import structs_parser +from wasm.codegen.helpers import structs class Generator: @@ -26,7 +26,7 @@ class Generator: def __init__(self): # Traverse the introspect dictionary to get the field # wrapper/bindings statements set up for each struct - self.structs_to_bind_data = structs_parser.generate_wasm_bindings( + self.structs_to_bind_data = structs.generate_wasm_bindings( constants.STRUCTS_TO_BIND ) @@ -38,7 +38,7 @@ class Generator: markers_and_content = [] # Sort by struct name by dependency to ensure deterministic output order - sorted_struct_names = structs_parser.sort_structs_by_dependency( + sorted_struct_names = structs.sort_structs_by_dependency( constants.STRUCTS_TO_BIND ) diff --git a/wasm/codegen/helpers/helpers_test.py b/wasm/codegen/helpers/helpers_test.py index 8b02c997..4cd495f1 100644 --- a/wasm/codegen/helpers/helpers_test.py +++ b/wasm/codegen/helpers/helpers_test.py @@ -21,10 +21,7 @@ from wasm.codegen.helpers import code_builder from wasm.codegen.helpers import common from wasm.codegen.helpers import constants from wasm.codegen.helpers import function_utils -from wasm.codegen.helpers import struct_constructor_code_builder -from wasm.codegen.helpers import struct_field_code_builder -from wasm.codegen.helpers import struct_field_handler -from wasm.codegen.helpers import structs_parser +from wasm.codegen.helpers import structs class CodeBuilderTest(absltest.TestCase): @@ -294,7 +291,7 @@ class FunctionUtilsTest(absltest.TestCase): class StructConstructorCodeBuilderTest(absltest.TestCase): def test_constructor_code_with_default_function(self): - wrapped_structs = structs_parser.generate_wasm_bindings(["mjLROpt"]) + wrapped_structs = structs.generate_wasm_bindings(["mjLROpt"]) self.assertEqual( wrapped_structs["mjLROpt"].wrapped_source, """ @@ -323,7 +320,7 @@ std::unique_ptr MjLROpt::copy() { ) def test_constructor_code_without_default_function(self): - wrapped_structs = structs_parser.generate_wasm_bindings(["mjsElement"]) + wrapped_structs = structs.generate_wasm_bindings(["mjsElement"]) self.assertEqual( wrapped_structs["mjsElement"].wrapped_source, """ @@ -343,11 +340,11 @@ std::unique_ptr MjsElement::copy() { ), doc="", ) - wrapped_field_data = struct_field_handler.StructFieldHandler( + wrapped_field_data = structs.StructFieldHandler( field_with_init, "MjsTexture" ).generate() self.assertEqual( - struct_constructor_code_builder.build_struct_source( + structs.build_struct_source( "mjsTexture", [wrapped_field_data], ), @@ -359,7 +356,7 @@ MjsTexture::~MjsTexture() {} def test_constructor_code_with_shallow_copy(self): self.assertEqual( - struct_constructor_code_builder.build_struct_source("mjvLight", []), + structs.build_struct_source("mjvLight", []), """MjvLight::MjvLight(mjvLight *ptr) : ptr_(ptr) {} MjvLight::MjvLight() : ptr_(new mjvLight) { owned_ = true; @@ -385,14 +382,14 @@ std::unique_ptr MjvLight::copy() { def test_build_struct_header_with_nested_wrappers(self): self.assertEqual( - struct_constructor_code_builder.build_struct_header("mjData", []), + structs.build_struct_header("mjData", []), "", ) def test_build_struct_header_basic_struct(self): self.assertEqual( - struct_constructor_code_builder.build_struct_header("mjLROpt", []), + structs.build_struct_header("mjLROpt", []), """ struct MjLROpt { MjLROpt(); @@ -420,7 +417,7 @@ class StructFieldCodeBuilderTest(absltest.TestCase): doc="number of geoms", ) self.assertEqual( - struct_field_code_builder.build_primitive_type_definition(field), + structs.build_primitive_type_definition(field), """ int ngeom() const { return ptr_->ngeom; @@ -441,7 +438,7 @@ void set_ngeom(int value) { array_extent=("ngeom", 4), ) self.assertEqual( - struct_field_code_builder.build_memory_view_definition( + structs.build_memory_view_definition( field, "ptr_->ngeom * 4", "ptr_->geom_rgba" ), """ @@ -460,7 +457,7 @@ emscripten::val geom_rgba() const { doc="rgba when material is omitted", ) self.assertEqual( - struct_field_code_builder.build_string_field_definition(field), + structs.build_string_field_definition(field), """ mjString string_field() const { return (ptr_ && ptr_->string_field) ? *(ptr_->string_field) : ""; @@ -482,9 +479,7 @@ void set_string_field(const mjString& value) { doc="", ) self.assertEqual( - struct_field_code_builder.build_mjvec_pointer_definition( - field, "mjDoubleVec" - ), + structs.build_mjvec_pointer_definition(field, "mjDoubleVec"), """ mjDoubleVec &vector_field() const { return *(ptr_->vector_field); @@ -500,9 +495,7 @@ mjDoubleVec &vector_field() const { doc="", ) self.assertEqual( - struct_field_code_builder.build_mjvec_pointer_definition( - field, "mjByteVec" - ), + structs.build_mjvec_pointer_definition(field, "mjByteVec"), """ std::vector &vector_field() const { return *(reinterpret_cast*>(ptr_->vector_field)); @@ -516,9 +509,7 @@ std::vector &vector_field() const { doc="number of geoms", ) self.assertEqual( - struct_field_code_builder.build_simple_property_binding( - field, "MjModel" - ), + structs.build_simple_property_binding(field, "MjModel"), '.property("ngeom", &MjModel::ngeom)', ) @@ -529,9 +520,7 @@ std::vector &vector_field() const { doc="", ) self.assertEqual( - struct_field_code_builder.build_simple_property_binding( - field, "MjModel", True - ), + structs.build_simple_property_binding(field, "MjModel", True), '.property("ngeom", &MjModel::ngeom, &MjModel::set_ngeom)', ) @@ -542,7 +531,7 @@ std::vector &vector_field() const { doc="", ) self.assertEqual( - struct_field_code_builder.build_simple_property_binding( + structs.build_simple_property_binding( field, "MjModel", add_setter=True, @@ -561,7 +550,7 @@ class StructFieldCodeBuilderTest(absltest.TestCase): doc="number of geoms", ) self.assertEqual( - struct_field_code_builder.build_primitive_type_definition(field), + structs.build_primitive_type_definition(field), """ int ngeom() const { return ptr_->ngeom; @@ -582,7 +571,7 @@ void set_ngeom(int value) { array_extent=("ngeom", 4), ) self.assertEqual( - struct_field_code_builder.build_memory_view_definition( + structs.build_memory_view_definition( field, "ptr_->ngeom * 4", "ptr_->geom_rgba" ), """ @@ -601,7 +590,7 @@ emscripten::val geom_rgba() const { doc="rgba when material is omitted", ) self.assertEqual( - struct_field_code_builder.build_string_field_definition(field), + structs.build_string_field_definition(field), """ mjString string_field() const { return (ptr_ && ptr_->string_field) ? *(ptr_->string_field) : ""; @@ -623,9 +612,7 @@ void set_string_field(const mjString& value) { doc="", ) self.assertEqual( - struct_field_code_builder.build_mjvec_pointer_definition( - field, "mjDoubleVec" - ), + structs.build_mjvec_pointer_definition(field, "mjDoubleVec"), """ mjDoubleVec &vector_field() const { return *(ptr_->vector_field); @@ -641,9 +628,7 @@ mjDoubleVec &vector_field() const { doc="", ) self.assertEqual( - struct_field_code_builder.build_mjvec_pointer_definition( - field, "mjByteVec" - ), + structs.build_mjvec_pointer_definition(field, "mjByteVec"), """ std::vector &vector_field() const { return *(reinterpret_cast*>(ptr_->vector_field)); @@ -657,9 +642,7 @@ std::vector &vector_field() const { doc="number of geoms", ) self.assertEqual( - struct_field_code_builder.build_simple_property_binding( - field, "MjModel" - ), + structs.build_simple_property_binding(field, "MjModel"), '.property("ngeom", &MjModel::ngeom)', ) @@ -670,9 +653,7 @@ std::vector &vector_field() const { doc="", ) self.assertEqual( - struct_field_code_builder.build_simple_property_binding( - field, "MjModel", True - ), + structs.build_simple_property_binding(field, "MjModel", True), '.property("ngeom", &MjModel::ngeom, &MjModel::set_ngeom)', ) @@ -683,7 +664,7 @@ std::vector &vector_field() const { doc="", ) self.assertEqual( - struct_field_code_builder.build_simple_property_binding( + structs.build_simple_property_binding( field, "MjModel", add_setter=True, @@ -703,9 +684,7 @@ class StructFieldHandlerTest(absltest.TestCase): doc="number of geoms", ) - field_handler_scalar = struct_field_handler.StructFieldHandler( - field_scalar, "MjModel" - ) + field_handler_scalar = structs.StructFieldHandler(field_scalar, "MjModel") wrapped_field_data = field_handler_scalar.generate() self.assertEqual( wrapped_field_data.definition, @@ -733,9 +712,7 @@ void set_ngeom(int value) { doc="rgba when material is omitted", array_extent=("ngeom", 4), ) - wrapped_field_data = struct_field_handler.StructFieldHandler( - field, "MjModel" - ).generate() + wrapped_field_data = structs.StructFieldHandler(field, "MjModel").generate() self.assertEqual( wrapped_field_data.definition, @@ -760,9 +737,7 @@ emscripten::val geom_rgba() const { ), doc="main buffer; all pointers point in it (nbuffer bytes)", ) - wrapped_field_data = struct_field_handler.StructFieldHandler( - field, "MjData" - ).generate() + wrapped_field_data = structs.StructFieldHandler(field, "MjData").generate() self.assertEqual( wrapped_field_data.definition, @@ -783,7 +758,7 @@ emscripten::val buffer() const { ), doc="", ) - wrapped_field_data = struct_field_handler.StructFieldHandler( + wrapped_field_data = structs.StructFieldHandler( field, "MjsTexture" ).generate() self.assertEqual(wrapped_field_data.definition, "MjsElement element;") @@ -807,7 +782,7 @@ emscripten::val buffer() const { ), doc="gravitational acceleration", ) - wrapped_field_data = struct_field_handler.StructFieldHandler( + wrapped_field_data = structs.StructFieldHandler( field, "MjOption" ).generate() @@ -834,9 +809,7 @@ emscripten::val gravity() const { ), doc="description", ) - wrapped_field_data = struct_field_handler.StructFieldHandler( - field, "MjModel" - ).generate() + wrapped_field_data = structs.StructFieldHandler(field, "MjModel").generate() self.assertEqual( wrapped_field_data.definition, """ @@ -853,52 +826,42 @@ emscripten::val multi_dim_array() const { def test_parse_array_extent(self): """Test that parse_array_extent handles various cases correctly.""" self.assertEqual( - struct_field_handler.parse_array_extent((1, 2), "MjModel", "geom_rgba"), + structs.parse_array_extent((1, 2), "MjModel", "geom_rgba"), "1 * 2", ) self.assertEqual( - struct_field_handler.parse_array_extent( - (1, "ngeom"), "MjModel", "geom_rgba" - ), + structs.parse_array_extent((1, "ngeom"), "MjModel", "geom_rgba"), "1 * ptr_->ngeom", ) self.assertEqual( - struct_field_handler.parse_array_extent( - (1, "mjConstant"), "MjModel", "geom_rgba" - ), + structs.parse_array_extent((1, "mjConstant"), "MjModel", "geom_rgba"), "1 * mjConstant", ) def test_resolve_extent(self): """Test that resolve_extent handles various cases correctly.""" # for integer just return the number - self.assertEqual( - struct_field_handler.resolve_extent(1, "MjModel", "geom_rgba"), "1" - ) + self.assertEqual(structs.resolve_extent(1, "MjModel", "geom_rgba"), "1") # for string that does not start with mj, it's a member of the struct self.assertEqual( - struct_field_handler.resolve_extent("ngeom", "MjModel", "geom_rgba"), + structs.resolve_extent("ngeom", "MjModel", "geom_rgba"), "ptr_->ngeom", ) # when it's MjData and the field is in MJDATA_SIZES, it should use ptr_-> self.assertEqual( - struct_field_handler.resolve_extent("size_value", "MjData", "efc_AR"), + structs.resolve_extent("size_value", "MjData", "efc_AR"), "ptr_->size_value", ) # when it's MjData and the field is not in MJDATA_SIZES, # it should use model-> self.assertEqual( - struct_field_handler.resolve_extent( - "size_value", "MjData", "data_field" - ), + structs.resolve_extent("size_value", "MjData", "data_field"), "model->size_value", ) self.assertEqual( - struct_field_handler.resolve_extent( - "mjConstant", "MjModel", "model_field" - ), + structs.resolve_extent("mjConstant", "MjModel", "model_field"), "mjConstant", ) diff --git a/wasm/codegen/helpers/struct_constructor_code_builder.py b/wasm/codegen/helpers/struct_constructor_code_builder.py deleted file mode 100644 index 03bd16aa..00000000 --- a/wasm/codegen/helpers/struct_constructor_code_builder.py +++ /dev/null @@ -1,232 +0,0 @@ -# Copyright 2025 DeepMind Technologies Limited -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""Code builder for struct constructor code.""" - -from typing import List, cast - -from introspect import ast_nodes -from introspect import structs as introspect_structs - -from wasm.codegen.helpers import code_builder -from wasm.codegen.helpers import common -from wasm.codegen.helpers import constants -from wasm.codegen.helpers import structs_wrappers_data - - -def _has_nested_wrapper_members(struct_info: ast_nodes.StructDecl) -> bool: - """Checks if the struct contains other wrapped structs as direct members.""" - for field in struct_info.fields: - struct_field = cast(ast_nodes.StructFieldDecl, field) - if isinstance(struct_field.type, ast_nodes.ValueType): - if struct_field.type.name in constants.STRUCTS_TO_BIND: - return True - if isinstance(struct_field.type, ast_nodes.ArrayType): - if isinstance(struct_field.type.inner_type, ast_nodes.ValueType): - if struct_field.type.inner_type.name in constants.STRUCTS_TO_BIND: - return True - if isinstance(struct_field.type, ast_nodes.PointerType): - if isinstance(struct_field.type.inner_type, ast_nodes.ValueType): - if struct_field.type.inner_type.name in constants.STRUCTS_TO_BIND: - return True - return False - - -def _build_struct_header_internal( - struct_name: str, - wrapped_fields: List[structs_wrappers_data.WrappedFieldData], - fields_with_init: List[structs_wrappers_data.WrappedFieldData], - is_mjs: bool = False, -): - """Builds the C++ header file code for a struct.""" - - shallow_copy = use_shallow_copy(wrapped_fields) - - wrapper_name = common.uppercase_first_letter(struct_name) - builder = code_builder.CodeBuilder() - with builder.block(f"struct {wrapper_name}"): - if not is_mjs: - builder.line(f"{wrapper_name}();") - builder.line(f"{wrapper_name}(const {wrapper_name} &);") - builder.line(f"{wrapper_name} &operator=(const {wrapper_name} &);") - - builder.line(f"explicit {wrapper_name}({struct_name} *ptr);") - builder.line(f"~{wrapper_name}();") - - if shallow_copy: - builder.line(f"std::unique_ptr<{wrapper_name}> copy();") - - for field in wrapped_fields: - if field.definition and field not in fields_with_init: - for line in field.definition.splitlines(): - builder.line(line) - - builder.line(f"{struct_name}* get() const {{ return ptr_; }}") - builder.line(f"void set({struct_name}* ptr) {{ ptr_ = ptr; }}") - - builder.newline() - builder.line("private:") - builder.line(f"{struct_name}* ptr_;") - if not is_mjs: - builder.line("bool owned_ = false;") - - if is_mjs and fields_with_init: - builder.newline() - builder.line("public:") - for field in fields_with_init: - if field.definition: - builder.line(f"{field.definition}") - return builder.to_string() + ";" - - -def _get_default_func_name(struct_name: str) -> str: - """Returns the default function name for the given struct.""" - 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"): - return f"mjs_default{struct_name.removeprefix('mjs')}" - elif struct_name.startswith("mjv"): - return f"mjv_default{struct_name.removeprefix('mjv')}" - else: - return f"mj_default{struct_name.removeprefix('mj')}" - - -def _find_fields_with_init( - wrapped_fields: List[structs_wrappers_data.WrappedFieldData], -) -> List[structs_wrappers_data.WrappedFieldData]: - """Finds the fields with initialization in the wrapped fields list.""" - fields_with_init = [] - for field in wrapped_fields: - if field.initialization: - fields_with_init.append(field) - return fields_with_init - - -def use_shallow_copy( - wrapped_fields: List[structs_wrappers_data.WrappedFieldData], -) -> bool: - """Returns true if the struct fields can be shallow copied.""" - for field in wrapped_fields: - if not field.is_primitive_or_fixed_size: - return False - return True - - -def build_struct_header( - struct_name: str, - wrapped_fields: List[structs_wrappers_data.WrappedFieldData], -): - """Builds the C++ header file code for a struct.""" - struct_info = introspect_structs.STRUCTS.get(struct_name) - - if struct_name.startswith("mjs"): - fields_with_init = _find_fields_with_init(wrapped_fields) - return _build_struct_header_internal( - struct_name, - wrapped_fields, - fields_with_init, - is_mjs=True, - ) - - if ( - ( - common.uppercase_first_letter(struct_name) - not in constants.HARDCODED_WRAPPER_STRUCTS - ) - and struct_info - and not _has_nested_wrapper_members(struct_info) - ): - return _build_struct_header_internal( - struct_name, wrapped_fields, [], is_mjs=False - ) - return "" - - -def build_struct_source( - struct_name: str, - wrapped_fields: List[structs_wrappers_data.WrappedFieldData], -): - """Builds the C++ .cc file code for a struct.""" - wrapper_name = common.uppercase_first_letter(struct_name) - is_mjs_struct = "Mjs" in wrapper_name - builder = code_builder.CodeBuilder() - - fields_with_init = _find_fields_with_init(wrapped_fields) - shallow_copy = use_shallow_copy(wrapped_fields) - mj_default_func = _get_default_func_name(struct_name) - - fields_init = "" - if fields_with_init: - fields_init = "".join( - field_with_init.initialization for field_with_init in fields_with_init - ) - # constructor passing native ptr - builder.line( - f"{wrapper_name}::{wrapper_name}({struct_name} *ptr) :" - f" ptr_(ptr){fields_init} {{}}" - ) - # constructor with default values - if not is_mjs_struct: - with builder.block( - f"{wrapper_name}::{wrapper_name}() : ptr_(new" - f" {struct_name}){fields_init}" - ): - builder.line("owned_ = true;") - if mj_default_func: - builder.line(f"{mj_default_func}(ptr_);") - # copy constructor - if shallow_copy and not is_mjs_struct: - with builder.block( - f"{wrapper_name}::{wrapper_name}(const {wrapper_name} &other)" - + (f" : {wrapper_name}()" if not is_mjs_struct else "") - ): - builder.line("*ptr_ = *other.get();") - if fields_with_init: - for field_with_init in fields_with_init: - if field_with_init.ptr_copy_reset is not None: - builder.line(field_with_init.ptr_copy_reset) - # assignment operator - with builder.block( - f"{wrapper_name}&" - f" {wrapper_name}::operator=(const" - f" {wrapper_name} &other)" - ): - with builder.block("if (this == &other)"): - builder.line("return *this;") - builder.line("*ptr_ = *other.get();") - if fields_with_init: - for field_with_init in fields_with_init: - if field_with_init.ptr_copy_reset is not None: - builder.line(field_with_init.ptr_copy_reset) - builder.line("return *this;") - # destructor - if is_mjs_struct: - builder.line(f"{wrapper_name}::~{wrapper_name}() {{}}") - else: - with builder.block(f"{wrapper_name}::~{wrapper_name}()"): - builder.line("if (owned_ && ptr_) delete ptr_;") - # copy function - if shallow_copy: - with builder.block( - f"std::unique_ptr<{wrapper_name}> {wrapper_name}::copy()" - ): - builder.line(f"return std::make_unique<{wrapper_name}>(*this);") - return builder.to_string() diff --git a/wasm/codegen/helpers/struct_field_code_builder.py b/wasm/codegen/helpers/struct_field_code_builder.py deleted file mode 100644 index 2a8ef8ae..00000000 --- a/wasm/codegen/helpers/struct_field_code_builder.py +++ /dev/null @@ -1,97 +0,0 @@ -# Copyright 2025 DeepMind Technologies Limited -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""Class to build the C++ code for a struct field wrapper.""" - -from introspect import ast_nodes -from wasm.codegen.helpers import code_builder - - -def build_primitive_type_definition(field: ast_nodes.StructFieldDecl) -> str: - """Builds the C++ 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() - # build getter for primitive type field - with builder.block(f"{field.type.name} {field.name}() const"): - builder.line(f"return ptr_->{field.name};") - # build setter for primitive type field - with builder.block(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.block(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 for a string type field wrapper.""" - builder = code_builder.CodeBuilder() - with builder.block(f"mjString {field.name}() const"): - builder.line( - f'return (ptr_ && ptr_->{field.name}) ? *(ptr_->{field.name}) : "";' - ) - with builder.block(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.block(f"{vector_type} &{field.name}() const"): - builder.line(f"return {ptr_field_expr};") - return builder.to_string() - - -def build_simple_property_binding( - field: ast_nodes.StructFieldDecl, - struct_wrapper_name: str, - add_setter: bool = False, - add_return_value_policy_as_ref: 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() diff --git a/wasm/codegen/helpers/struct_field_handler.py b/wasm/codegen/helpers/struct_field_handler.py deleted file mode 100644 index 933af915..00000000 --- a/wasm/codegen/helpers/struct_field_handler.py +++ /dev/null @@ -1,351 +0,0 @@ -# Copyright 2025 DeepMind Technologies Limited -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""Class to handle the different struct field types, and provide the c++ code for the wrappers and bindings.""" - -import math -from typing import Tuple, Union, cast -from introspect import ast_nodes -from wasm.codegen.helpers import common -from wasm.codegen.helpers import constants -from wasm.codegen.helpers import struct_field_code_builder -from wasm.codegen.helpers import structs_wrappers_data - -debug_print = common.debug_print - - -class StructFieldHandler: - """Class to handle the different struct field types, and provide the c++ code for the definitions and bindings.""" - - 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 = ( - struct_field_code_builder.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) -> structs_wrappers_data.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) -> structs_wrappers_data.WrappedFieldData: - """Handles the generation of C++ definition and binding code for primitive fields.""" - return structs_wrappers_data.WrappedFieldData( - definition=( - struct_field_code_builder.build_primitive_type_definition( - self.field - ) - ), - binding=struct_field_code_builder.build_simple_property_binding( - self.field, - self.struct_wrapper_name, - add_setter=True, - add_return_value_policy_as_ref=True, - ), - is_primitive_or_fixed_size=True, - ) - - def _handle_pointer(self) -> structs_wrappers_data.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}" - ) - field_type: ast_nodes.PointerType = self.field.type - inner_type_name = ( - field_type.inner_type.name - if isinstance(field_type.inner_type, ast_nodes.ValueType) - else "" - ) - ptr_field_expr = f"ptr_->{self.field.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(): - # 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 - ) - elif inner_type_name == "mjString": - return structs_wrappers_data.WrappedFieldData( - definition=struct_field_code_builder.build_string_field_definition( - self.field - ), - binding=struct_field_code_builder.build_simple_property_binding( - self.field, - self.struct_wrapper_name, - add_setter=True, - add_return_value_policy_as_ref=True, - ), - ) - elif inner_type_name.startswith("mj") and inner_type_name.endswith("Vec"): - return structs_wrappers_data.WrappedFieldData( - definition=struct_field_code_builder.build_mjvec_pointer_definition( - self.field, inner_type_name - ), - binding=struct_field_code_builder.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" - ) - - if ( - inner_type_name.startswith("mj") - and inner_type_name not in constants.PRIMITIVE_TYPES - ): - debug_print( - f"\tcomplex pointer type: needs manual wrapper: {self.field.name}" - ) - # 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 structs_wrappers_data.WrappedFieldData( - definition=f"{wrapper_field_name} {self.field.name};", - binding=struct_field_code_builder.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: - debug_print( - "\tcomplex pointer type with array extent: needs manual wrapper:" - f" {self.field.name}" - ) - return self._get_manual_definition(comment_type="complex pointer field") - - return structs_wrappers_data.WrappedFieldData( - definition=( - struct_field_code_builder.build_memory_view_definition( - self.field, array_size_str, ptr_field_expr - ) - ), - binding=self.simple_property_binding, - ) - - def _handle_array(self) -> structs_wrappers_data.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}" - ) - 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 structs_wrappers_data.WrappedFieldData( - definition=( - struct_field_code_builder.build_memory_view_definition( - self.field, str(size), ptr_expr - ) - ), - binding=self.simple_property_binding, - is_primitive_or_fixed_size=True, - ) - elif inner_type.name.startswith("mj") and not inner_type.name.startswith( - "mjt" - ): - debug_print(f"\tarray to vector wrapper needed: {self.field.name}") - return self._get_manual_definition(comment_type="array field") - - debug_print(f"\tNOT IMPLEMENTED ARRAY field: {self.field.name}") - return structs_wrappers_data.WrappedFieldData( - definition=( - f"// TODO: NOT IMPLEMENTED ARRAY wrapper for {self.field.name}" - ), - binding=f"// TODO: NOT IMPLEMENTED ARRAY binding for {self.field.name}", - ) - - def _handle_mj_struct(self) -> structs_wrappers_data.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 structs_wrappers_data.WrappedFieldData( - definition=definition, - binding=struct_field_code_builder.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, - ) - return self._get_manual_definition(comment_type="struct field") - - def _handle_anonymous_struct(self) -> structs_wrappers_data.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 structs_wrappers_data.WrappedFieldData( - binding=struct_field_code_builder.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, - ) - return self._get_manual_definition(comment_type="anonymous struct field") - - def _undefined(self) -> structs_wrappers_data.WrappedFieldData: - """This function adds a TODO comment for fields that are not handled by this class yet.""" - return structs_wrappers_data.WrappedFieldData( - definition=f"// TODO: UNDEFINED definition for {self.field.name}", - binding=f"// TODO: UNDEFINED binding for {self.field.name}", - ) - - def _get_manual_definition( - self, comment_type: str = "" - ) -> structs_wrappers_data.WrappedFieldData: - """Helper method to generate a comment as a definition for manually added fields.""" - if self.field.name in self.manually_added_fields: - return structs_wrappers_data.WrappedFieldData( - definition=( - f"// {comment_type} is defined manually. {self.field.name}" - ), - binding=self.simple_property_binding, - ) - - return structs_wrappers_data.WrappedFieldData( - definition=( - f"// TODO: Define {comment_type} manually for {self.field.name}" - ), - binding=f"// TODO: {self.simple_property_binding}", - ) - - -def parse_array_extent( - extents: Tuple[Union[str, int], ...], wrapper_name: str, field_name: str -) -> str: - """Parses the array extent of a field, returning a string representing the resolved extents.""" - if not extents: - return "" - return " * ".join( - resolve_extent(extent, wrapper_name, field_name) for extent in extents - ) - - -def resolve_extent( - extent: Union[str, int], wrapper_name: str, field_name: str -) -> str: - """Resolves the extent of an array, handling integers and references to other struct fields. - - Args: - extent: The extent to resolve, can be an int or a string referencing a - field. - wrapper_name: The name of the struct wrapper. - field_name: The name of the field being processed. - - Returns: - A string representing the resolved extent, either as a number or a field - reference. - """ - if isinstance(extent, int): - return str(extent) - # if starts with mj, it's a mujoco constant, - # so we don't need to get a parent struct ptr - if extent.startswith("mj"): - return str(extent) - if wrapper_name == "MjData" and field_name not in constants.MJDATA_SIZES: - var_name = "model" - else: - var_name = "ptr_" - return f"{var_name}->{extent}" diff --git a/wasm/codegen/helpers/structs.py b/wasm/codegen/helpers/structs.py new file mode 100644 index 00000000..befdd1b5 --- /dev/null +++ b/wasm/codegen/helpers/structs.py @@ -0,0 +1,817 @@ +# Copyright 2025 DeepMind Technologies Limited +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Parser for MuJoCo structs.""" + +import collections +import dataclasses +import math +from typing import Dict, List, Tuple, Union, cast + +from introspect import ast_nodes +from introspect import structs + +from wasm.codegen.helpers import code_builder +from wasm.codegen.helpers import common +from wasm.codegen.helpers import constants + + +debug_print = common.debug_print + +introspect_structs = structs.STRUCTS + + +@dataclasses.dataclass +class WrappedFieldData: + """Data class for struct field definition and binding.""" + + # Line for struct field binding + binding: str + + # Line for struct field definition + definition: str | None = None + + # Initialization code for fields that require it + initialization: str | None = None + + # Statement to reset the inner pointer when copying the field + ptr_copy_reset: str | None = None + + # Whether the field is a primitive or fixed size + is_primitive_or_fixed_size: bool = False + + +@dataclasses.dataclass +class WrappedStructData: + """Data class for struct wrapper definition and binding.""" + + # Name of wrapper struct + wrap_name: str + + # List of WrappedFieldData for this struct + wrapped_fields: List[WrappedFieldData] + + # Struct header code + wrapped_header: str + + # Struct source code + wrapped_source: str + + # Whether to use shallow copy for this struct + use_shallow_copy: bool = True + + +def build_primitive_type_definition(field: ast_nodes.StructFieldDecl) -> str: + """Builds the C++ 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() + # build getter for primitive type field + with builder.block(f"{field.type.name} {field.name}() const"): + builder.line(f"return ptr_->{field.name};") + # build setter for primitive type field + with builder.block(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.block(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 for a string type field wrapper.""" + builder = code_builder.CodeBuilder() + with builder.block(f"mjString {field.name}() const"): + builder.line( + f'return (ptr_ && ptr_->{field.name}) ? *(ptr_->{field.name}) : "";' + ) + with builder.block(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.block(f"{vector_type} &{field.name}() const"): + builder.line(f"return {ptr_field_expr};") + return builder.to_string() + + +def build_simple_property_binding( + field: ast_nodes.StructFieldDecl, + struct_wrapper_name: str, + add_setter: bool = False, + add_return_value_policy_as_ref: 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() + + +class StructFieldHandler: + """Class to handle the different struct field types, and provide the c++ code for the definitions and bindings.""" + + 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.""" + return WrappedFieldData( + definition=(build_primitive_type_definition(self.field)), + binding=build_simple_property_binding( + self.field, + self.struct_wrapper_name, + add_setter=True, + add_return_value_policy_as_ref=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}" + ) + field_type: ast_nodes.PointerType = self.field.type + inner_type_name = ( + field_type.inner_type.name + if isinstance(field_type.inner_type, ast_nodes.ValueType) + else "" + ) + ptr_field_expr = f"ptr_->{self.field.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(): + # 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 + ) + elif inner_type_name == "mjString": + return WrappedFieldData( + definition=build_string_field_definition(self.field), + binding=build_simple_property_binding( + self.field, + self.struct_wrapper_name, + add_setter=True, + add_return_value_policy_as_ref=True, + ), + ) + elif inner_type_name.startswith("mj") and inner_type_name.endswith("Vec"): + return WrappedFieldData( + definition=build_mjvec_pointer_definition( + self.field, inner_type_name + ), + 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" + ) + + if ( + inner_type_name.startswith("mj") + and inner_type_name not in constants.PRIMITIVE_TYPES + ): + debug_print( + f"\tcomplex pointer type: needs manual wrapper: {self.field.name}" + ) + # 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};", + 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: + debug_print( + "\tcomplex pointer type with array extent: needs manual wrapper:" + f" {self.field.name}" + ) + return self._get_manual_definition(comment_type="complex pointer field") + + return WrappedFieldData( + 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}" + ) + 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) + ), + binding=self.simple_property_binding, + is_primitive_or_fixed_size=True, + ) + elif inner_type.name.startswith("mj") and not inner_type.name.startswith( + "mjt" + ): + debug_print(f"\tarray to vector wrapper needed: {self.field.name}") + return self._get_manual_definition(comment_type="array field") + + debug_print(f"\tNOT IMPLEMENTED ARRAY field: {self.field.name}") + return WrappedFieldData( + definition=( + f"// TODO: NOT IMPLEMENTED ARRAY wrapper for {self.field.name}" + ), + 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, + 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, + ) + 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, + ), + 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, + ) + 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}", + ) + + 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( + definition=( + f"// {comment_type} is defined manually. {self.field.name}" + ), + binding=self.simple_property_binding, + ) + + return WrappedFieldData( + definition=( + f"// TODO: Define {comment_type} manually for {self.field.name}" + ), + binding=f"// TODO: {self.simple_property_binding}", + ) + + +def _has_nested_wrapper_members(struct_info: ast_nodes.StructDecl) -> bool: + """Checks if the struct contains other wrapped structs as direct members.""" + for field in struct_info.fields: + struct_field = cast(ast_nodes.StructFieldDecl, field) + if isinstance(struct_field.type, ast_nodes.ValueType): + if struct_field.type.name in constants.STRUCTS_TO_BIND: + return True + if isinstance(struct_field.type, ast_nodes.ArrayType): + if isinstance(struct_field.type.inner_type, ast_nodes.ValueType): + if struct_field.type.inner_type.name in constants.STRUCTS_TO_BIND: + return True + if isinstance(struct_field.type, ast_nodes.PointerType): + if isinstance(struct_field.type.inner_type, ast_nodes.ValueType): + if struct_field.type.inner_type.name in constants.STRUCTS_TO_BIND: + return True + return False + + +def _build_struct_header_internal( + struct_name: str, + wrapped_fields: List[WrappedFieldData], + fields_with_init: List[WrappedFieldData], + is_mjs: bool = False, +): + """Builds the C++ header file code for a struct.""" + + shallow_copy = use_shallow_copy(wrapped_fields) + + wrapper_name = common.uppercase_first_letter(struct_name) + builder = code_builder.CodeBuilder() + with builder.block(f"struct {wrapper_name}"): + if not is_mjs: + builder.line(f"{wrapper_name}();") + builder.line(f"{wrapper_name}(const {wrapper_name} &);") + builder.line(f"{wrapper_name} &operator=(const {wrapper_name} &);") + + builder.line(f"explicit {wrapper_name}({struct_name} *ptr);") + builder.line(f"~{wrapper_name}();") + + if shallow_copy: + builder.line(f"std::unique_ptr<{wrapper_name}> copy();") + + for field in wrapped_fields: + if field.definition and field not in fields_with_init: + for line in field.definition.splitlines(): + builder.line(line) + + builder.line(f"{struct_name}* get() const {{ return ptr_; }}") + builder.line(f"void set({struct_name}* ptr) {{ ptr_ = ptr; }}") + + builder.newline() + builder.line("private:") + builder.line(f"{struct_name}* ptr_;") + if not is_mjs: + builder.line("bool owned_ = false;") + + if is_mjs and fields_with_init: + builder.newline() + builder.line("public:") + for field in fields_with_init: + if field.definition: + builder.line(f"{field.definition}") + return builder.to_string() + ";" + + +def _get_default_func_name(struct_name: str) -> str: + """Returns the default function name for the given struct.""" + 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"): + return f"mjs_default{struct_name.removeprefix('mjs')}" + elif struct_name.startswith("mjv"): + return f"mjv_default{struct_name.removeprefix('mjv')}" + else: + return f"mj_default{struct_name.removeprefix('mj')}" + + +def _find_fields_with_init( + wrapped_fields: List[WrappedFieldData], +) -> List[WrappedFieldData]: + """Finds the fields with initialization in the wrapped fields list.""" + fields_with_init = [] + for field in wrapped_fields: + if field.initialization: + fields_with_init.append(field) + return fields_with_init + + +def use_shallow_copy( + wrapped_fields: List[WrappedFieldData], +) -> bool: + """Returns true if the struct fields can be shallow copied.""" + for field in wrapped_fields: + if not field.is_primitive_or_fixed_size: + return False + return True + + +def build_struct_header( + struct_name: str, + wrapped_fields: List[WrappedFieldData], +): + """Builds the C++ header file code for a struct.""" + struct_info = introspect_structs.get(struct_name) + + if struct_name.startswith("mjs"): + fields_with_init = _find_fields_with_init(wrapped_fields) + return _build_struct_header_internal( + struct_name, + wrapped_fields, + fields_with_init, + is_mjs=True, + ) + + if ( + ( + common.uppercase_first_letter(struct_name) + not in constants.HARDCODED_WRAPPER_STRUCTS + ) + and struct_info + and not _has_nested_wrapper_members(struct_info) + ): + return _build_struct_header_internal( + struct_name, wrapped_fields, [], is_mjs=False + ) + return "" + + +def build_struct_source( + struct_name: str, + wrapped_fields: List[WrappedFieldData], +): + """Builds the C++ .cc file code for a struct.""" + wrapper_name = common.uppercase_first_letter(struct_name) + is_mjs_struct = "Mjs" in wrapper_name + builder = code_builder.CodeBuilder() + + fields_with_init = _find_fields_with_init(wrapped_fields) + shallow_copy = use_shallow_copy(wrapped_fields) + mj_default_func = _get_default_func_name(struct_name) + + fields_init = "" + if fields_with_init: + fields_init = "".join( + field_with_init.initialization for field_with_init in fields_with_init + ) + # constructor passing native ptr + builder.line( + f"{wrapper_name}::{wrapper_name}({struct_name} *ptr) :" + f" ptr_(ptr){fields_init} {{}}" + ) + # constructor with default values + if not is_mjs_struct: + with builder.block( + f"{wrapper_name}::{wrapper_name}() : ptr_(new" + f" {struct_name}){fields_init}" + ): + builder.line("owned_ = true;") + if mj_default_func: + builder.line(f"{mj_default_func}(ptr_);") + # copy constructor + if shallow_copy and not is_mjs_struct: + with builder.block( + f"{wrapper_name}::{wrapper_name}(const {wrapper_name} &other)" + + (f" : {wrapper_name}()" if not is_mjs_struct else "") + ): + builder.line("*ptr_ = *other.get();") + if fields_with_init: + for field_with_init in fields_with_init: + if field_with_init.ptr_copy_reset is not None: + builder.line(field_with_init.ptr_copy_reset) + # assignment operator + with builder.block( + f"{wrapper_name}&" + f" {wrapper_name}::operator=(const" + f" {wrapper_name} &other)" + ): + with builder.block("if (this == &other)"): + builder.line("return *this;") + builder.line("*ptr_ = *other.get();") + if fields_with_init: + for field_with_init in fields_with_init: + if field_with_init.ptr_copy_reset is not None: + builder.line(field_with_init.ptr_copy_reset) + builder.line("return *this;") + # destructor + if is_mjs_struct: + builder.line(f"{wrapper_name}::~{wrapper_name}() {{}}") + else: + with builder.block(f"{wrapper_name}::~{wrapper_name}()"): + builder.line("if (owned_ && ptr_) delete ptr_;") + # copy function + if shallow_copy: + with builder.block( + f"std::unique_ptr<{wrapper_name}> {wrapper_name}::copy()" + ): + builder.line(f"return std::make_unique<{wrapper_name}>(*this);") + return builder.to_string() + + +def parse_array_extent( + extents: Tuple[Union[str, int], ...], wrapper_name: str, field_name: str +) -> str: + """Parses the array extent of a field, returning a string representing the resolved extents.""" + if not extents: + return "" + return " * ".join( + resolve_extent(extent, wrapper_name, field_name) for extent in extents + ) + + +def resolve_extent( + extent: Union[str, int], wrapper_name: str, field_name: str +) -> str: + """Resolves the extent of an array, handling integers and references to other struct fields. + + Args: + extent: The extent to resolve, can be an int or a string referencing a + field. + wrapper_name: The name of the struct wrapper. + field_name: The name of the field being processed. + + Returns: + A string representing the resolved extent, either as a number or a field + reference. + """ + if isinstance(extent, int): + return str(extent) + # if starts with mj, it's a mujoco constant, + # so we don't need to get a parent struct ptr + if extent.startswith("mj"): + return str(extent) + if wrapper_name == "MjData" and field_name not in constants.MJDATA_SIZES: + var_name = "model" + else: + var_name = "ptr_" + return f"{var_name}->{extent}" + + +def generate_wasm_bindings( + structs_to_bind: List[str], +) -> Dict[str, WrappedStructData]: + """Generates WASM bindings for MuJoCo structs.""" + + wrapped_structs: Dict[str, WrappedStructData] = {} + for struct_name in structs_to_bind: + wrapped_name = common.uppercase_first_letter(struct_name) + + if struct_name in introspect_structs: + struct_fields = introspect_structs[struct_name].fields + elif struct_name in constants.ANONYMOUS_STRUCTS: + anonymous_struct = _get_anonymous_struct_field(struct_name) + if not anonymous_struct or not isinstance( + anonymous_struct.type, ast_nodes.AnonymousStructDecl + ): + raise RuntimeError(f"Anonymous struct not found: {struct_name}") + struct_fields = anonymous_struct.type.fields + else: + raise RuntimeError(f"Struct not found: {struct_name}") + + debug_print(f"Wrapping struct: {struct_name}") + + wrapped_fields: List[WrappedFieldData] = [] + for field in struct_fields: + wrapped_field = StructFieldHandler(field, wrapped_name).generate() + wrapped_fields.append(wrapped_field) + + wrapped_header = build_struct_header( + struct_name, + wrapped_fields, + ) + wrapped_source = build_struct_source( + struct_name, + wrapped_fields, + ) + wrap_data = WrappedStructData( + wrap_name=wrapped_name, + wrapped_fields=wrapped_fields, + wrapped_header=wrapped_header, + wrapped_source=wrapped_source, + use_shallow_copy=use_shallow_copy(wrapped_fields), + ) + + wrapped_structs[struct_name] = wrap_data + + return wrapped_structs + + +def _get_anonymous_struct_field( + anonymous_structs_key: str, +) -> ast_nodes.StructFieldDecl | None: + """Looks up the given key in the anonymous_structs dict and generates bindings for its fields.""" + info = constants.ANONYMOUS_STRUCTS[anonymous_structs_key] + parent_decl = introspect_structs[info["parent"]] + target_field = next( + ( + f + for f in parent_decl.fields + if hasattr(f, "name") + and f.name == info["field_name"] + and hasattr(f, "type") + and isinstance(f.type, ast_nodes.AnonymousStructDecl) + ), + None, + ) + return target_field + + +def _get_field_struct_type(field_type): + """Extracts the base struct name if the field type is a struct or pointer to a struct.""" + if isinstance(field_type, ast_nodes.ValueType): + return field_type.name + if isinstance(field_type, ast_nodes.PointerType): + if isinstance(field_type.inner_type, ast_nodes.ValueType): + return field_type.inner_type.name + return None + + +def sort_structs_by_dependency(struct_names: List[str]) -> List[str]: + """Sorts structs based on their field dependencies using topological sort. + + Structs with no dependencies on other structs in the list come first. + If struct A has a field of type struct B, B must come before A in the + sorted list. + + Args: + struct_names: A list of struct names to sort. + + Returns: + A new list of struct names sorted by dependency. + + Raises: + RuntimeError: If a cyclic dependency is detected. + """ + adj = collections.defaultdict(list) + in_degree = collections.defaultdict(int) + struct_set = set(struct_names) + sorted_struct_names = sorted(struct_names) + + for struct_name in sorted_struct_names: + if struct_name not in introspect_structs: + # Skip anonymous or other structs not in the main introspect map + continue + + struct_decl = introspect_structs[struct_name] + for field in struct_decl.fields: + if isinstance(field, ast_nodes.AnonymousStructDecl): + continue + + field_type_name = _get_field_struct_type(field.type) + if ( + field_type_name + and field_type_name != struct_name + and field_type_name in struct_set + ): + if struct_name not in adj[field_type_name]: + adj[field_type_name].append(struct_name) + in_degree[struct_name] += 1 + + queue = collections.deque( + [name for name in sorted_struct_names if in_degree[name] == 0] + ) + sorted_list = [] + + while queue: + u = queue.popleft() + sorted_list.append(u) + for v in adj[u]: + in_degree[v] -= 1 + if in_degree[v] == 0: + queue.append(v) + + if len(sorted_list) == len(struct_names): + return sorted_list + else: + remaining = set(struct_names) - set(sorted_list) + raise RuntimeError( + "Cycle detected in struct dependencies, involving: " + f"{', '.join(sorted(list(remaining)))}" + ) diff --git a/wasm/codegen/helpers/structs_parser.py b/wasm/codegen/helpers/structs_parser.py deleted file mode 100644 index 1303dcd2..00000000 --- a/wasm/codegen/helpers/structs_parser.py +++ /dev/null @@ -1,179 +0,0 @@ -# Copyright 2025 DeepMind Technologies Limited -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""Parser for MuJoCo structs.""" - -import collections -from typing import Dict, List - -from introspect import ast_nodes -from introspect import structs - -from wasm.codegen.helpers import common -from wasm.codegen.helpers import constants -from wasm.codegen.helpers import struct_constructor_code_builder -from wasm.codegen.helpers import struct_field_handler -from wasm.codegen.helpers import structs_wrappers_data - - -debug_print = common.debug_print - -introspect_structs = structs.STRUCTS - - -def generate_wasm_bindings( - structs_to_bind: List[str], -) -> Dict[str, structs_wrappers_data.WrappedStructData]: - """Generates WASM bindings for MuJoCo structs.""" - - wrapped_structs: Dict[str, structs_wrappers_data.WrappedStructData] = {} - for struct_name in structs_to_bind: - wrapped_name = common.uppercase_first_letter(struct_name) - - if struct_name in introspect_structs: - struct_fields = introspect_structs[struct_name].fields - elif struct_name in constants.ANONYMOUS_STRUCTS: - anonymous_struct = _get_anonymous_struct_field(struct_name) - if not anonymous_struct or not isinstance( - anonymous_struct.type, ast_nodes.AnonymousStructDecl - ): - raise RuntimeError(f"Anonymous struct not found: {struct_name}") - struct_fields = anonymous_struct.type.fields - else: - raise RuntimeError(f"Struct not found: {struct_name}") - - debug_print(f"Wrapping struct: {struct_name}") - - wrapped_fields: List[structs_wrappers_data.WrappedFieldData] = [] - for field in struct_fields: - wrapped_field = struct_field_handler.StructFieldHandler( - field, wrapped_name - ).generate() - wrapped_fields.append(wrapped_field) - - wrapped_header = struct_constructor_code_builder.build_struct_header( - struct_name, - wrapped_fields, - ) - wrapped_source = struct_constructor_code_builder.build_struct_source( - struct_name, - wrapped_fields, - ) - wrap_data = structs_wrappers_data.WrappedStructData( - wrap_name=wrapped_name, - wrapped_fields=wrapped_fields, - wrapped_header=wrapped_header, - wrapped_source=wrapped_source, - use_shallow_copy=struct_constructor_code_builder.use_shallow_copy( - wrapped_fields - ), - ) - - wrapped_structs[struct_name] = wrap_data - - return wrapped_structs - - -def _get_anonymous_struct_field( - anonymous_structs_key: str, -) -> ast_nodes.StructFieldDecl | None: - """Looks up the given key in the anonymous_structs dict and generates bindings for its fields.""" - info = constants.ANONYMOUS_STRUCTS[anonymous_structs_key] - parent_decl = introspect_structs[info["parent"]] - target_field = next( - ( - f - for f in parent_decl.fields - if hasattr(f, "name") - and f.name == info["field_name"] - and hasattr(f, "type") - and isinstance(f.type, ast_nodes.AnonymousStructDecl) - ), - None, - ) - return target_field - - -def _get_field_struct_type(field_type): - """Extracts the base struct name if the field type is a struct or pointer to a struct.""" - if isinstance(field_type, ast_nodes.ValueType): - return field_type.name - if isinstance(field_type, ast_nodes.PointerType): - if isinstance(field_type.inner_type, ast_nodes.ValueType): - return field_type.inner_type.name - return None - - -def sort_structs_by_dependency(struct_names: List[str]) -> List[str]: - """Sorts structs based on their field dependencies using topological sort. - - Structs with no dependencies on other structs in the list come first. - If struct A has a field of type struct B, B must come before A in the - sorted list. - - Args: - struct_names: A list of struct names to sort. - - Returns: - A new list of struct names sorted by dependency. - - Raises: - RuntimeError: If a cyclic dependency is detected. - """ - adj = collections.defaultdict(list) - in_degree = collections.defaultdict(int) - struct_set = set(struct_names) - sorted_struct_names = sorted(struct_names) - - for struct_name in sorted_struct_names: - if struct_name not in introspect_structs: - # Skip anonymous or other structs not in the main introspect map - continue - - struct_decl = introspect_structs[struct_name] - for field in struct_decl.fields: - if isinstance(field, ast_nodes.AnonymousStructDecl): - continue - - field_type_name = _get_field_struct_type(field.type) - if ( - field_type_name - and field_type_name != struct_name - and field_type_name in struct_set - ): - if struct_name not in adj[field_type_name]: - adj[field_type_name].append(struct_name) - in_degree[struct_name] += 1 - - queue = collections.deque( - [name for name in sorted_struct_names if in_degree[name] == 0] - ) - sorted_list = [] - - while queue: - u = queue.popleft() - sorted_list.append(u) - for v in adj[u]: - in_degree[v] -= 1 - if in_degree[v] == 0: - queue.append(v) - - if len(sorted_list) == len(struct_names): - return sorted_list - else: - remaining = set(struct_names) - set(sorted_list) - raise RuntimeError( - "Cycle detected in struct dependencies, involving: " - f"{', '.join(sorted(list(remaining)))}" - ) diff --git a/wasm/codegen/helpers/structs_wrappers_data.py b/wasm/codegen/helpers/structs_wrappers_data.py deleted file mode 100644 index 6d4ae717..00000000 --- a/wasm/codegen/helpers/structs_wrappers_data.py +++ /dev/null @@ -1,61 +0,0 @@ -# Copyright 2025 DeepMind Technologies Limited -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""Classes used by the parser to generate the structs wrappers and bindings statements.""" - -import dataclasses -from typing import Dict, List - -from wasm.codegen.helpers import common - - -@dataclasses.dataclass -class WrappedFieldData: - """Data class for struct field definition and binding.""" - - # Line for struct field binding - binding: str - - # Line for struct field definition - definition: str | None = None - - # Initialization code for fields that require it - initialization: str | None = None - - # Statement to reset the inner pointer when copying the field - ptr_copy_reset: str | None = None - - # Whether the field is a primitive or fixed size - is_primitive_or_fixed_size: bool = False - - -@dataclasses.dataclass -class WrappedStructData: - """Data class for struct wrapper definition and binding.""" - - # Name of wrapper struct - wrap_name: str - - # List of WrappedFieldData for this struct - wrapped_fields: List[WrappedFieldData] - - # Struct header code - wrapped_header: str - - # Struct source code - wrapped_source: str - - # Whether to use shallow copy for this struct - use_shallow_copy: bool = True -