diff --git a/wasm/README.md b/wasm/README.md index e35cc367..6be76a4c 100644 --- a/wasm/README.md +++ b/wasm/README.md @@ -148,7 +148,7 @@ bounds checking and nice error reporting. > run the following command to re-create the test: > > ```sh - > PYTHONPATH=python/mujoco python3 -m wasm.codegen.enums_test_generator + > PYTHONPATH=python/mujoco python3 -m wasm.codegen.tests.enums_test_generator > ``` 2. **JavaScript API benchmark tests.** diff --git a/wasm/codegen/generators/binding_builder.py b/wasm/codegen/generators/binding_builder.py new file mode 100644 index 00000000..322c29a6 --- /dev/null +++ b/wasm/codegen/generators/binding_builder.py @@ -0,0 +1,71 @@ +# 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. + +from introspect import ast_nodes +from introspect import enums as introspect_enums +from introspect import functions as introspect_functions + +from wasm.codegen.generators import common +from wasm.codegen.generators import constants +from wasm.codegen.generators import enums +from wasm.codegen.generators import functions +from wasm.codegen.generators import structs + + +class BindingBuilder: + """Builds WASM bindings for MuJoCo.""" + + def __init__( + self, + template_path_cc: str, + ): + with open(template_path_cc, "r") as f: + self.content_cc = f.readlines() + + self.markers_and_content = [] + + def set_enums(self): + """Generates and sets the enum bindings.""" + generator = enums.Generator(introspect_enums.ENUMS) + self.markers_and_content += generator.generate() + return self + + def set_structs(self): + """Generates and sets the struct bindings.""" + generator = structs.Generator() + self.markers_and_content += generator.generate() + return self + + def set_functions(self): + """Generates and sets the function wrappers and bindings.""" + + functions_to_bind: dict[str, ast_nodes.FunctionDecl] = {} + for name, func in introspect_functions.FUNCTIONS.items(): + if not functions.is_excluded_function_name(name): + if name not in constants.BOUNDCHECK_FUNCS: + functions_to_bind[name] = func + + generator = functions.Generator(functions_to_bind) + self.markers_and_content += generator.generate() + return self + + def to_string(self) -> str: + for marker, content in self.markers_and_content: + self.content_cc = common.replace_lines_containing_marker( + self.content_cc, marker, content + ) + return "".join(self.content_cc) + + def build(self, generated_path_cc: str): + common.write_to_file(generated_path_cc, self.to_string()) diff --git a/wasm/codegen/helpers/code_builder.py b/wasm/codegen/generators/code_builder.py similarity index 100% rename from wasm/codegen/helpers/code_builder.py rename to wasm/codegen/generators/code_builder.py diff --git a/wasm/codegen/helpers/common.py b/wasm/codegen/generators/common.py similarity index 100% rename from wasm/codegen/helpers/common.py rename to wasm/codegen/generators/common.py diff --git a/wasm/codegen/helpers/constants.py b/wasm/codegen/generators/constants.py similarity index 100% rename from wasm/codegen/helpers/constants.py rename to wasm/codegen/generators/constants.py diff --git a/wasm/codegen/generators/enums.py b/wasm/codegen/generators/enums.py index b9f3a7a4..d98aa25e 100644 --- a/wasm/codegen/generators/enums.py +++ b/wasm/codegen/generators/enums.py @@ -14,11 +14,11 @@ """Generates Embind bindings for MuJoCo enums.""" -from typing import Mapping, Optional +from typing import Mapping from introspect import ast_nodes -from wasm.codegen.helpers import code_builder +from wasm.codegen.generators import code_builder class Generator: diff --git a/wasm/codegen/generators/functions.py b/wasm/codegen/generators/functions.py index bdfbcaa7..f4da342a 100644 --- a/wasm/codegen/generators/functions.py +++ b/wasm/codegen/generators/functions.py @@ -12,14 +12,340 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Generates Embind bindings for MuJoCo functions.""" +"""Helper functions for processing and generating bindings for MuJoCo functions.""" -from typing import List, Mapping +from typing import List, Mapping, Set, Tuple, cast from introspect import ast_nodes -from wasm.codegen.helpers import code_builder -from wasm.codegen.helpers import functions as function_utils +from wasm.codegen.generators import code_builder +from wasm.codegen.generators import common +from wasm.codegen.generators import constants + + +PRIMITIVE_TYPES = constants.PRIMITIVE_TYPES +uppercase_first_letter = common.uppercase_first_letter + + +def param_is_primitive_value(param: ast_nodes.FunctionParameterDecl) -> bool: + """Checks if param is a primitive value type.""" + if isinstance(param.type, ast_nodes.ValueType): + return param.type.name in PRIMITIVE_TYPES + return False + + +def param_is_pointer_to_primitive_value( + param: ast_nodes.FunctionParameterDecl, +) -> bool: + """Checks if param is a pointer to a primitive value.""" + return ( + isinstance(param.type, ast_nodes.PointerType) + or isinstance(param.type, ast_nodes.ArrayType) + ) and ( + isinstance(param.type.inner_type, ast_nodes.ValueType) + and param.type.inner_type.name in PRIMITIVE_TYPES + ) + + +def param_is_pointer_to_struct(param: ast_nodes.FunctionParameterDecl) -> bool: + """Checks if param is a pointer to a struct.""" + return ( + isinstance(param.type, ast_nodes.PointerType) + or isinstance(param.type, ast_nodes.ArrayType) + ) and ( + isinstance(param.type.inner_type, ast_nodes.ValueType) + and param.type.inner_type.name not in PRIMITIVE_TYPES + ) + + +def return_is_value_of_type( + func: ast_nodes.FunctionDecl, allowed_types: Set[str] +) -> bool: + """Checks if func returns an allowed value type.""" + return ( + isinstance(func.return_type, ast_nodes.ValueType) + and func.return_type.name in allowed_types + ) + + +def return_is_pointer_to_struct(func: ast_nodes.FunctionDecl) -> bool: + """Checks if func returns a pointer to a struct.""" + return ( + isinstance(func.return_type, ast_nodes.PointerType) + and isinstance(func.return_type.inner_type, ast_nodes.ValueType) + and func.return_type.inner_type.name not in PRIMITIVE_TYPES + ) + + +def return_is_pointer_to_primitive(func: ast_nodes.FunctionDecl) -> bool: + """Checks if func returns a pointer to a primitive value.""" + return ( + isinstance(func.return_type, ast_nodes.PointerType) + and isinstance(func.return_type.inner_type, ast_nodes.ValueType) + and func.return_type.inner_type.name in PRIMITIVE_TYPES + ) + + +def get_const_qualifier(func: ast_nodes.FunctionDecl) -> str: + """Returns the const qualifier of func's return type.""" + if ( + isinstance(func.return_type, ast_nodes.PointerType) + and isinstance(func.return_type.inner_type, ast_nodes.ValueType) + and func.return_type.inner_type.is_const + ): + return "const " + return "" + + +def should_be_wrapped(func: ast_nodes.FunctionDecl) -> bool: + """Checks if a MuJoCo function needs a wrapper function.""" + return ( + return_is_pointer_to_primitive(func) + or return_is_pointer_to_struct(func) + or any( + param_is_pointer_to_primitive_value(param) + or isinstance(param.type, ast_nodes.ArrayType) + or param_is_pointer_to_struct(param) + for param in func.parameters + ) + ) + + +def generate_function_wrapper(func: ast_nodes.FunctionDecl) -> str: + """Generates C++ code for a wrapper function.""" + + params_unpack_statements = get_params_unpack_statements(func.parameters) + wrapper_params_list = get_params_string(func.parameters) + not_nullable_params = get_params_notnullable(func.parameters) + wrapper_params = ", ".join(wrapper_params_list) + ret_type = get_compatible_return_type(func) + + builder = code_builder.CodeBuilder() + with builder.function(f"{ret_type} {func.name}_wrapper({wrapper_params})"): + invoker_params_list = get_params_string_maybe_with_conversion( + func.parameters + ) + invoker_params_str = ", ".join(invoker_params_list) + invoker_call = f"{func.name}({invoker_params_str})" + invoker_statement = get_compatible_return_call(func, invoker_call) + for p in not_nullable_params: + builder.line(f"CHECK_VAL({p});") + for unpack_statement in params_unpack_statements: + builder.line(unpack_statement) + builder.line(f"{invoker_statement};") + + return builder.to_string() + + +def get_params_notnullable( + ast_params: Tuple[ast_nodes.FunctionParameterDecl, ...], +) -> List[str]: + """Generates list of param names for checking if they aren't null/undefined.""" + + not_nullable_params = [] + for p in ast_params: + if ( + isinstance(p.type, (ast_nodes.PointerType, ast_nodes.ArrayType)) + and isinstance(p.type.inner_type, ast_nodes.ValueType) + # We only check for char because others are checked in the unpacker + # and we don't want to check twice. + and p.type.inner_type.name == "char" + and not p.nullable + ): + not_nullable_params.append(p.name) + return not_nullable_params + + +def get_params_unpack_statements( + ast_params: Tuple[ast_nodes.FunctionParameterDecl, ...], +) -> List[str]: + """Generates C++ statements to unpack JS values for pointer/array parameters.""" + + params_unpack_statements = [] + for p in ast_params: + if ( + isinstance(p.type, (ast_nodes.PointerType, ast_nodes.ArrayType)) + and isinstance(p.type.inner_type, ast_nodes.ValueType) + and p.type.inner_type.name in PRIMITIVE_TYPES + ): + if p.type.inner_type.name == "char": + # param is Javascript string + continue + + if p.type.inner_type.is_const: + # param is Javascript number[] + params_unpack_statements.append( + f"UNPACK_ARRAY({p.type.inner_type.name}, {p.name});" + ) + else: + # param is TypedArray or a WasmBuffer + params_unpack_statements.append( + f"UNPACK_VALUE({p.type.inner_type.name}, {p.name});" + ) + return params_unpack_statements + + +def get_params_string( + parameters: Tuple[ast_nodes.FunctionParameterDecl, ...] +) -> List[str]: + """Generates a list of C++ parameter declarations as strings.""" + + result = [] + for p in parameters: + if ( + isinstance(p.type, ast_nodes.PointerType) + and isinstance(p.type.inner_type, ast_nodes.ValueType) + and p.type.inner_type.name not in PRIMITIVE_TYPES + ): + # Pointer to struct parameters + const_qualifier = "const " if p.type.inner_type.is_const else "" + result.append( + f"{const_qualifier}{uppercase_first_letter(p.type.inner_type.name)}&" + f" {p.name}" + ) + elif ( + isinstance(p.type, ast_nodes.ValueType) + and p.type.name in PRIMITIVE_TYPES + ): + # Primitive value parameters + const_qualifier = "const " if p.type.is_const else "" + result.append(f"{const_qualifier}{p.type} {p.name}") + elif ( + isinstance(p.type, (ast_nodes.PointerType, ast_nodes.ArrayType)) + and isinstance(p.type.inner_type, ast_nodes.ValueType) + and p.type.inner_type.name in PRIMITIVE_TYPES + ): + # Pointer to primitive value parameters or arrays + if p.type.inner_type.name == "char": + if p.nullable: + result.append(f"const NullableString& {p.name}") + else: + result.append(f"const String& {p.name}") + elif ( + p.type.inner_type.name + in ["int", "float", "double", "mjtNum", "mjtByte"] + and p.type.inner_type.is_const + ): + result.append(f"const NumberArray& {p.name}") + else: + result.append(f"const val& {p.name}") + else: + # This case should ideally not be reached if AST is well-formed + # and types are categorized by the helper booleans correctly. + raise TypeError( + "Unable to generate param string. Unhandled parameter type:" + f" {p.type} for param '{p.name}'" + ) + return result + + +def get_params_string_maybe_with_conversion( + ast_params: Tuple[ast_nodes.FunctionParameterDecl, ...], +) -> List[str]: + """Generates C++ expressions for passing compatible params from JS to MuJoCo C-API functions.""" + + native_params = [] + for p in ast_params: + if param_is_pointer_to_struct(p): + native_params.append(f"{p.name}.get()") + elif param_is_primitive_value(p): + native_params.append(p.name) + elif ( + isinstance(p.type, (ast_nodes.PointerType, ast_nodes.ArrayType)) + and isinstance(p.type.inner_type, ast_nodes.ValueType) + and p.type.inner_type.name in PRIMITIVE_TYPES + and p.type.inner_type.name != "char" + ): + native_params.append(f"{p.name}_.data()") + elif ( + isinstance(p.type, (ast_nodes.PointerType, ast_nodes.ArrayType)) + and isinstance(p.type.inner_type, ast_nodes.ValueType) + and p.type.inner_type.name == "char" + ): + const_qualifier = "const " if p.type.inner_type.is_const else "" + native_params.append( + f"{p.name}.as<{const_qualifier}std::string>().data()" + ) + else: + raise TypeError( + f"Unhandled parameter type for conversion: {p.type} for param" + f" '{p.name}'" + ) + return native_params + + +def get_compatible_return_call( + func: ast_nodes.FunctionDecl, invoker: str +) -> str: + """Generates embind compatible return value conversion.""" + + if return_is_value_of_type(func, {"void"}): + return invoker + if isinstance(func.return_type, ast_nodes.PointerType) and isinstance( + func.return_type.inner_type, ast_nodes.ValueType + ): + if func.return_type.inner_type.name == "char": + return f"return std::string({invoker})" + elif func.return_type.inner_type.name == "mjString": + return f"return *{invoker}" + if return_is_pointer_to_struct(func): + return get_converted_struct_to_class(func, invoker) + if return_is_value_of_type(func, PRIMITIVE_TYPES): + return f"return {invoker}" + raise RuntimeError( + "Failed to calculate return value conversion for function" + f" {func.name} that returns '{func.return_type}'" + ) + + +def get_compatible_return_type(func: ast_nodes.FunctionDecl) -> str: + """Creates embind compatible return type.""" + + if ( + isinstance(func.return_type, ast_nodes.PointerType) + and isinstance(func.return_type.inner_type, ast_nodes.ValueType) + and func.return_type.inner_type.name in ["char", "mjString"] + ): + return "std::string" + if ( + isinstance(func.return_type, ast_nodes.PointerType) + and isinstance(func.return_type.inner_type, ast_nodes.ValueType) + and func.return_type.inner_type.name not in PRIMITIVE_TYPES + ): + const_qualifier = get_const_qualifier(func) + return f"""{const_qualifier}std::optional<{uppercase_first_letter(func.return_type.inner_type.name)}>""" + if ( + isinstance(func.return_type, ast_nodes.ValueType) + and func.return_type.name in PRIMITIVE_TYPES + ): + return f"{func.return_type.name}" + return "val" + + +def get_converted_struct_to_class( + func: ast_nodes.FunctionDecl, invoker: str +) -> str: + """Generates a C++ function invocation for a struct return-type function.""" + + const_qualifier = get_const_qualifier(func) + return_type = cast(ast_nodes.PointerType, func.return_type) + struct_name = cast(ast_nodes.ValueType, return_type.inner_type).name + class_constructor = uppercase_first_letter(struct_name) + return_str = f"{class_constructor}(result)" + return f"""{const_qualifier}{struct_name}* result = {invoker}; + if (result == nullptr) {{ + return std::nullopt; + }} + return {return_str}""" + + +def is_excluded_function_name(func_name: str) -> bool: + """Checks if a function name should be excluded from direct binding.""" + return ( + func_name.startswith("mjr_") + or func_name.startswith("mjui_") + or func_name in constants.SKIPPED_FUNCTIONS + ) class Generator: @@ -30,7 +356,7 @@ class Generator: self.wrapper_bind_functions: List[ast_nodes.FunctionDecl] = [] for func in functions.values(): - if function_utils.should_be_wrapped(func): + if should_be_wrapped(func): self.wrapper_bind_functions.append(func) else: self.direct_bind_functions.append(func) @@ -40,7 +366,7 @@ class Generator: code = [] for func in self.wrapper_bind_functions: - wrapper_code = function_utils.generate_function_wrapper(func) + wrapper_code = generate_function_wrapper(func) code.append(wrapper_code) return "\n\n".join(code) diff --git a/wasm/codegen/generators/generators_test.py b/wasm/codegen/generators/generators_test.py deleted file mode 100644 index bf61e90f..00000000 --- a/wasm/codegen/generators/generators_test.py +++ /dev/null @@ -1,106 +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. - -from absl.testing import absltest -from introspect import ast_nodes - -from wasm.codegen.generators import enums -from wasm.codegen.generators import functions - - -class EnumsGeneratorTest(absltest.TestCase): - - def test_generate_enum_bindings(self): - - generator = enums.Generator({ - "TestEnum": ast_nodes.EnumDecl( - name="TestEnum", - declname="enum TestEnum_", - values={"FIRST_VAL": 0, "SECOND_VAL": 1, "THIRD_VAL": 2}, - ), - "AnotherEnum": ast_nodes.EnumDecl( - name="AnotherEnum", - declname="enum AnotherEnum_", - values={"ALPHA": 100, "BETA": 200}, - ), - "EmptyEnum": ast_nodes.EnumDecl( - name="EmptyEnum", - declname="enum EmptyEnum_", - values={}, - ), - }) - - expected_code = """ enum_("TestEnum") - .value("FIRST_VAL", FIRST_VAL) - .value("SECOND_VAL", SECOND_VAL) - .value("THIRD_VAL", THIRD_VAL); - - enum_("AnotherEnum") - .value("ALPHA", ALPHA) - .value("BETA", BETA); - - enum_("EmptyEnum");""" - - markers_and_content = generator.generate() - actual_code = "\n\n".join(markers_and_content[0][1]) - - self.assertEqual(actual_code, expected_code) - - -class FunctionsGeneratorTest(absltest.TestCase): - - def setUp(self): - super().setUp() - self.generator = functions.Generator({}) - self.int_type = ast_nodes.ValueType(name="int") - - def test_generate_function_binding_simple_case(self): - func_simple_void = ast_nodes.FunctionDecl( - name="do_nothing", - return_type=ast_nodes.ValueType(name="void"), - parameters=tuple(), - doc="doc", - ) - self.assertEqual( - self.generator._generate_function_binding(func_simple_void), - 'function("do_nothing", &do_nothing);\n', - ) - - def test_generate_direct_bindable_functions_simple_filter(self): - direct_bind = ast_nodes.FunctionDecl( - name="direct_bind", - return_type=self.int_type, - parameters=( - ast_nodes.FunctionParameterDecl(name="val", type=self.int_type), - ), - doc="doc", - ) - needs_wrap = ast_nodes.FunctionDecl( - name="needs_wrap", - return_type=ast_nodes.PointerType(inner_type=self.int_type), - parameters=tuple(), - doc="doc", - ) - self.generator = functions.Generator({ - "direct1": direct_bind, - "wrapped1": needs_wrap, - }) - - generated_code = self.generator._generate_direct_bindable_functions() - self.assertIn('function("direct_bind", &direct_bind);\n', generated_code) - self.assertNotIn('function("needs_wrap", &needs_wrap);\n', generated_code) - - -if __name__ == "__main__": - absltest.main() diff --git a/wasm/codegen/generators/structs.py b/wasm/codegen/generators/structs.py index 8833c118..bb3d7a78 100644 --- a/wasm/codegen/generators/structs.py +++ b/wasm/codegen/generators/structs.py @@ -14,8 +14,701 @@ """Generates Embind bindings for MuJoCo structs.""" -from wasm.codegen.helpers import constants -from wasm.codegen.helpers import structs +import collections +import dataclasses +import math +from typing import Dict, List, Tuple, Union, cast + +from introspect import ast_nodes +from introspect import structs as introspect_structs + +from wasm.codegen.generators import code_builder +from wasm.codegen.generators import common +from wasm.codegen.generators import constants + + +@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 = "" + + # Initialization code for fields that require it + initialization: str = "" + + # Statement to reset the inner pointer when copying the field + ptr_copy_reset: str = "" + + # Whether the field is a primitive or fixed size + is_primitive_or_fixed_size: bool = False + + # 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 = "" + + +@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 + + # Struct bindings code + bindings: str = "" + + +def _simple_property_binding( + field: ast_nodes.StructFieldDecl, + struct_wrapper_name: str, + setter: bool = False, + reference: bool = False, +) -> str: + """Builds the C++ code for a simple property binding.""" + 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})' + + +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) + + 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( + 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, + ) + + 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, + ) + + 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, + ) + + elif isinstance(f.type, ast_nodes.PointerType): + + inner_type_name = ( + f.type.inner_type.name + if isinstance(f.type.inner_type, ast_nodes.ValueType) + else "" + ) + ptr_field_expr = f"ptr_->{f.name}" + array_size_str = "" + + 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[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=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=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() + ): + ptr_field = cast(ast_nodes.PointerType, f.type) + wrapper_field_name = common.uppercase_first_letter( + cast(ast_nodes.ValueType, ptr_field.inner_type).name + ) + return WrappedFieldData( + 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})", + ) + + 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 WrappedFieldData( + definition=builder.to_string(), + typename=_get_field_struct_type(f.type), + binding=_simple_property_binding(f, w), + ) + + # SHOULD NOT OCCUR + print("Error: field {f.name} not properly handled") + return WrappedFieldData( + definition=f"// Error: field {f.name} not properly handled.", + typename=_get_field_struct_type(f.type), + binding=f"// Error: field {f.name} not properly handled.", + ) + + +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: + member_type = cast(ast_nodes.StructFieldDecl, field).type + if isinstance(member_type, (ast_nodes.ArrayType, ast_nodes.PointerType)): + member_type = member_type.inner_type + if isinstance(member_type, ast_nodes.ValueType): + return member_type.name in constants.STRUCTS_TO_BIND + 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.""" + s = struct_name + w = common.uppercase_first_letter(s) + + shallow_copy = use_shallow_copy(wrapped_fields) + + builder = code_builder.CodeBuilder() + with builder.struct(f"{w}"): + builder.line(f"explicit {w}({s} *ptr);") + builder.line(f"~{w}();") + + if not is_mjs: + builder.line(f"{w}();") + + if shallow_copy and not is_mjs: + builder.line(f"{w}(const {w} &);") + builder.line(f"{w} &operator=(const {w} &);") + builder.line(f"std::unique_ptr<{w}> copy();") + + builder.line(f"{s}* get() const;") + builder.line(f"void set({s}* ptr);") + + 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.private() + builder.line(f"{s}* ptr_;") + if not is_mjs: + builder.line("bool owned_ = false;") + + if is_mjs and fields_with_init: + builder.public() + for field in fields_with_init: + if field.definition: + builder.line(f"{field.definition}") + + return builder.to_string() + ";" + + +def _default_function_statement(struct_name: str) -> str: + """Returns the default function name for the given struct.""" + if struct_name == "mjvGeom": + f = "mjv_initGeom(ptr_, mjGEOM_NONE, nullptr, nullptr, nullptr, nullptr);" + return f + elif struct_name in constants.ANONYMOUS_STRUCTS.keys(): + return "" + elif struct_name in constants.NO_DEFAULT_CONSTRUCTORS: + return "" + elif struct_name.startswith("mjs"): + return f"mjs_default{struct_name.removeprefix('mjs')}(ptr_);" + elif struct_name.startswith("mjv"): + return f"mjv_default{struct_name.removeprefix('mjv')}(ptr_);" + elif struct_name.startswith("mj"): + return f"mj_default{struct_name.removeprefix('mj')}(ptr_);" + + return "" + + +def _delete_ptr_statement(struct_name: str) -> str: + """Returns the delete function name for the given struct.""" + if struct_name == "mjVFS": + return "mj_deleteVFS(ptr_);" + + return "delete ptr_;" + + +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.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, + ) + + is_anonymous_struct = struct_name in constants.ANONYMOUS_STRUCTS.keys() + is_hardcoded_wrapper_struct = ( + common.uppercase_first_letter(struct_name) in constants.MANUAL_STRUCTS + ) + + if ( + not is_hardcoded_wrapper_struct + and struct_info + and not _has_nested_wrapper_members(struct_info) + or is_anonymous_struct + ): + 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.""" + # These structs require specific function calls for creation and/or deletion + # which, for now, are hardcoded in the template file. + if struct_name in [ + "mjData", + "mjModel", + "mjvScene", + "mjSpec", + ]: + return "" + + s = struct_name + w = common.uppercase_first_letter(s) + is_mjs = w.startswith("Mjs") + + fields_with_init = _find_fields_with_init(wrapped_fields) + shallow_copy = use_shallow_copy(wrapped_fields) + + fields_init = "" + if fields_with_init: + fields_init = "".join( + field_with_init.initialization for field_with_init in fields_with_init + ) + + builder = code_builder.CodeBuilder() + + # constructor passing native ptr + with builder.function(f"{w}::{w}({s} *ptr) : ptr_(ptr){fields_init}"): + pass + + # destructor + with builder.function(f"{w}::~{w}()"): + if not is_mjs: + with builder.block("if (owned_ && ptr_)"): + delete_ptr = _delete_ptr_statement(s) + builder.line(delete_ptr) + + if not is_mjs: + # default constructor + with builder.function(f"{w}::{w}() : ptr_(new {s}){fields_init}"): + builder.line("owned_ = true;") + default_func = _default_function_statement(s) + if default_func: + builder.line(default_func) + + if shallow_copy and not is_mjs: + # copy constructor + with builder.function(f"{w}::{w}(const {w} &other) : {w}()"): + builder.line("*ptr_ = *other.get();") + 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.function(f"{w}& {w}::operator=(const {w} &other)"): + with builder.block("if (this == &other)"): + builder.line("return *this;") + builder.line("*ptr_ = *other.get();") + 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;") + + # explicit copy function + with builder.function(f"std::unique_ptr<{w}> {w}::copy()"): + builder.line(f"return std::make_unique<{w}>(*this);") + + # C struct getter/setter + with builder.function(f"{s}* {w}::get() const"): + builder.line("return ptr_;") + with builder.function(f"void {w}::set({s}* ptr)"): + builder.line("ptr_ = ptr;") + + return builder.to_string() + + +def _build_struct_bindings( + struct_name: str, + wrapped_fields: List[WrappedFieldData], +): + """Builds the C++ bindings for a struct.""" + w = common.uppercase_first_letter(struct_name) + is_mjs = w.startswith("Mjs") + + builder = code_builder.CodeBuilder() + with builder.block( + header_line=f'emscripten::class_<{w}>("{w}")', braces=False + ): + if w == "MjData": + builder.line(".constructor()") + builder.line(".constructor()") + elif w == "MjModel": + builder.line( + '.class_function("loadFromXML", &loadFromXML, take_ownership())' + ) + builder.line(".constructor()") + elif w == "MjSpec": + builder.line(".constructor()") + elif w == "MjvScene": + builder.line(".constructor()") + builder.line(".constructor<>()") + elif not is_mjs: + builder.line(".constructor<>()") + + shallow_copy = use_shallow_copy(wrapped_fields) + if shallow_copy and not is_mjs: + builder.line(f'.function("copy", &{w}::copy, take_ownership())') + + for field in wrapped_fields[:-1]: + if field.binding: + builder.line(field.binding) + if wrapped_fields: + builder.line(f"{wrapped_fields[-1].binding};") + + 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: + s = struct_name + w = common.uppercase_first_letter(s) + + if s in introspect_structs.STRUCTS: + struct_fields = introspect_structs.STRUCTS[s].fields + elif s in constants.ANONYMOUS_STRUCTS: + anonymous_struct = _get_anonymous_struct_field(s) + if not anonymous_struct or not isinstance( + anonymous_struct.type, ast_nodes.AnonymousStructDecl + ): + raise RuntimeError(f"Anonymous struct not found: {s}") + struct_fields = anonymous_struct.type.fields + else: + raise RuntimeError(f"Struct not found: {s}") + + wrapped_fields: List[WrappedFieldData] = [] + for field in struct_fields: + wrapped_fields.append(_generate_field_data(field, w)) + + wrap_data = WrappedStructData( + wrap_name=w, + wrapped_fields=wrapped_fields, + wrapped_header=build_struct_header(s, wrapped_fields), + wrapped_source=build_struct_source(s, wrapped_fields), + bindings=_build_struct_bindings(s, wrapped_fields), + ) + + wrapped_structs[s] = 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.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_wrappers: dict[str, WrappedStructData], +) -> List[str]: + """Sorts structs based on their field dependencies using topological sort. + + Structs with no dependencies on other structs in the list come first. + Struct A has a dependency on struct B if struct A has a field where the + underlying_type is B. Note that this definition is stricter than the C++ + struct dependency criterion where forward declarations can be used to + eliminate dependencies A and B if A only has a pointer to B. + + Args: + struct_wrappers: A dictionary mapping struct names to their + WrappedStructData. + + 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_names = struct_wrappers.keys() + struct_set = set(struct_names) + sorted_struct_names = sorted(struct_names) + + for struct_name in sorted_struct_names: + for field in struct_wrappers[struct_name].wrapped_fields: + + field_type_name = field.typename + 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)))}" + ) class Generator: @@ -26,7 +719,7 @@ class Generator: # Traverse the introspect dictionary to get the field # wrapper/bindings statements set up for each struct - self.structs_to_bind_data = structs.generate_wasm_bindings( + self.structs_to_bind_data = generate_wasm_bindings( constants.STRUCTS_TO_BIND ) @@ -34,16 +727,12 @@ class Generator: markers_and_content = [] # Sort by struct name by dependency to ensure deterministic output order - sorted_struct_names = structs.sort_structs_by_dependency( - self.structs_to_bind_data - ) + sorted_struct_names = sort_structs_by_dependency(self.structs_to_bind_data) for struct_name in sorted_struct_names: struct_data = self.structs_to_bind_data[struct_name] if struct_data.wrapped_header: - autogenned_struct_definitions.append( - struct_data.wrapped_header + "\n" - ) + autogenned_struct_definitions.append(struct_data.wrapped_header + "\n") else: markers_and_content.append(( f"// INSERT-GENERATED-{struct_data.wrap_name}-DEFINITIONS", @@ -62,9 +751,7 @@ class Generator: for struct_name in sorted_struct_names: struct_data = self.structs_to_bind_data[struct_name] if struct_data.wrapped_source: - autogenned_struct_source.append( - struct_data.wrapped_source + "\n" - ) + autogenned_struct_source.append(struct_data.wrapped_source + "\n") autogenned_struct_bindings.append(struct_data.bindings) markers_and_content.append(( diff --git a/wasm/codegen/helpers/functions.py b/wasm/codegen/helpers/functions.py deleted file mode 100644 index c0e528fc..00000000 --- a/wasm/codegen/helpers/functions.py +++ /dev/null @@ -1,348 +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. - -"""Helper functions for processing and generating bindings for MuJoCo functions.""" - -from typing import List, Set, Tuple, cast - -from introspect import ast_nodes - -from wasm.codegen.helpers import code_builder -from wasm.codegen.helpers import common -from wasm.codegen.helpers import constants - - -PRIMITIVE_TYPES = constants.PRIMITIVE_TYPES -uppercase_first_letter = common.uppercase_first_letter - - -def param_is_primitive_value(param: ast_nodes.FunctionParameterDecl) -> bool: - """Checks if param is a primitive value type.""" - if isinstance(param.type, ast_nodes.ValueType): - return param.type.name in PRIMITIVE_TYPES - return False - - -def param_is_pointer_to_primitive_value( - param: ast_nodes.FunctionParameterDecl, -) -> bool: - """Checks if param is a pointer to a primitive value.""" - return ( - isinstance(param.type, ast_nodes.PointerType) - or isinstance(param.type, ast_nodes.ArrayType) - ) and ( - isinstance(param.type.inner_type, ast_nodes.ValueType) - and param.type.inner_type.name in PRIMITIVE_TYPES - ) - - -def param_is_pointer_to_struct(param: ast_nodes.FunctionParameterDecl) -> bool: - """Checks if param is a pointer to a struct.""" - return ( - isinstance(param.type, ast_nodes.PointerType) - or isinstance(param.type, ast_nodes.ArrayType) - ) and ( - isinstance(param.type.inner_type, ast_nodes.ValueType) - and param.type.inner_type.name not in PRIMITIVE_TYPES - ) - - -def return_is_value_of_type( - func: ast_nodes.FunctionDecl, allowed_types: Set[str] -) -> bool: - """Checks if func returns an allowed value type.""" - return ( - isinstance(func.return_type, ast_nodes.ValueType) - and func.return_type.name in allowed_types - ) - - -def return_is_pointer_to_struct(func: ast_nodes.FunctionDecl) -> bool: - """Checks if func returns a pointer to a struct.""" - return ( - isinstance(func.return_type, ast_nodes.PointerType) - and isinstance(func.return_type.inner_type, ast_nodes.ValueType) - and func.return_type.inner_type.name not in PRIMITIVE_TYPES - ) - - -def return_is_pointer_to_primitive(func: ast_nodes.FunctionDecl) -> bool: - """Checks if func returns a pointer to a primitive value.""" - return ( - isinstance(func.return_type, ast_nodes.PointerType) - and isinstance(func.return_type.inner_type, ast_nodes.ValueType) - and func.return_type.inner_type.name in PRIMITIVE_TYPES - ) - - -def get_const_qualifier(func: ast_nodes.FunctionDecl) -> str: - """Returns the const qualifier of func's return type.""" - if ( - isinstance(func.return_type, ast_nodes.PointerType) - and isinstance(func.return_type.inner_type, ast_nodes.ValueType) - and func.return_type.inner_type.is_const - ): - return "const " - return "" - - -def should_be_wrapped(func: ast_nodes.FunctionDecl) -> bool: - """Checks if a MuJoCo function needs a wrapper function.""" - return ( - return_is_pointer_to_primitive(func) - or return_is_pointer_to_struct(func) - or any( - param_is_pointer_to_primitive_value(param) - or isinstance(param.type, ast_nodes.ArrayType) - or param_is_pointer_to_struct(param) - for param in func.parameters - ) - ) - - -def generate_function_wrapper(func: ast_nodes.FunctionDecl) -> str: - """Generates C++ code for a wrapper function.""" - - params_unpack_statements = get_params_unpack_statements(func.parameters) - wrapper_params_list = get_params_string(func.parameters) - not_nullable_params = get_params_notnullable(func.parameters) - wrapper_params = ", ".join(wrapper_params_list) - ret_type = get_compatible_return_type(func) - - builder = code_builder.CodeBuilder() - with builder.function(f"{ret_type} {func.name}_wrapper({wrapper_params})"): - invoker_params_list = get_params_string_maybe_with_conversion( - func.parameters - ) - invoker_params_str = ", ".join(invoker_params_list) - invoker_call = f"{func.name}({invoker_params_str})" - invoker_statement = get_compatible_return_call(func, invoker_call) - for p in not_nullable_params: - builder.line(f"CHECK_VAL({p});") - for unpack_statement in params_unpack_statements: - builder.line(unpack_statement) - builder.line(f"{invoker_statement};") - - return builder.to_string() - - -def get_params_notnullable( - ast_params: Tuple[ast_nodes.FunctionParameterDecl, ...], -) -> List[str]: - """Generates list of param names for checking if they aren't null/undefined.""" - - not_nullable_params = [] - for p in ast_params: - if ( - isinstance(p.type, (ast_nodes.PointerType, ast_nodes.ArrayType)) - and isinstance(p.type.inner_type, ast_nodes.ValueType) - # We only check for char because others are checked in the unpacker - # and we don't want to check twice. - and p.type.inner_type.name == "char" - and not p.nullable - ): - not_nullable_params.append(p.name) - return not_nullable_params - - -def get_params_unpack_statements( - ast_params: Tuple[ast_nodes.FunctionParameterDecl, ...], -) -> List[str]: - """Generates C++ statements to unpack JS values for pointer/array parameters.""" - - params_unpack_statements = [] - for p in ast_params: - if ( - isinstance(p.type, (ast_nodes.PointerType, ast_nodes.ArrayType)) - and isinstance(p.type.inner_type, ast_nodes.ValueType) - and p.type.inner_type.name in PRIMITIVE_TYPES - ): - if p.type.inner_type.name == "char": - # param is Javascript string - continue - - if p.type.inner_type.is_const: - # param is Javascript number[] - params_unpack_statements.append( - f"UNPACK_ARRAY({p.type.inner_type.name}, {p.name});" - ) - else: - # param is TypedArray or a WasmBuffer - params_unpack_statements.append( - f"UNPACK_VALUE({p.type.inner_type.name}, {p.name});" - ) - return params_unpack_statements - - -def get_params_string( - parameters: Tuple[ast_nodes.FunctionParameterDecl, ...] -) -> List[str]: - """Generates a list of C++ parameter declarations as strings.""" - - result = [] - for p in parameters: - if ( - isinstance(p.type, ast_nodes.PointerType) - and isinstance(p.type.inner_type, ast_nodes.ValueType) - and p.type.inner_type.name not in PRIMITIVE_TYPES - ): - # Pointer to struct parameters - const_qualifier = "const " if p.type.inner_type.is_const else "" - result.append( - f"{const_qualifier}{uppercase_first_letter(p.type.inner_type.name)}&" - f" {p.name}" - ) - elif ( - isinstance(p.type, ast_nodes.ValueType) - and p.type.name in PRIMITIVE_TYPES - ): - # Primitive value parameters - const_qualifier = "const " if p.type.is_const else "" - result.append(f"{const_qualifier}{p.type} {p.name}") - elif ( - isinstance(p.type, (ast_nodes.PointerType, ast_nodes.ArrayType)) - and isinstance(p.type.inner_type, ast_nodes.ValueType) - and p.type.inner_type.name in PRIMITIVE_TYPES - ): - # Pointer to primitive value parameters or arrays - if p.type.inner_type.name == "char": - if p.nullable: - result.append(f"const NullableString& {p.name}") - else: - result.append(f"const String& {p.name}") - elif ( - p.type.inner_type.name - in ["int", "float", "double", "mjtNum", "mjtByte"] - and p.type.inner_type.is_const - ): - result.append(f"const NumberArray& {p.name}") - else: - result.append(f"const val& {p.name}") - else: - # This case should ideally not be reached if AST is well-formed - # and types are categorized by the helper booleans correctly. - raise TypeError( - "Unable to generate param string. Unhandled parameter type:" - f" {p.type} for param '{p.name}'" - ) - return result - - -def get_params_string_maybe_with_conversion( - ast_params: Tuple[ast_nodes.FunctionParameterDecl, ...], -) -> List[str]: - """Generates C++ expressions for passing compatible params from JS to MuJoCo C-API functions.""" - - native_params = [] - for p in ast_params: - if param_is_pointer_to_struct(p): - native_params.append(f"{p.name}.get()") - elif param_is_primitive_value(p): - native_params.append(p.name) - elif ( - isinstance(p.type, (ast_nodes.PointerType, ast_nodes.ArrayType)) - and isinstance(p.type.inner_type, ast_nodes.ValueType) - and p.type.inner_type.name in PRIMITIVE_TYPES - and p.type.inner_type.name != "char" - ): - native_params.append(f"{p.name}_.data()") - elif ( - isinstance(p.type, (ast_nodes.PointerType, ast_nodes.ArrayType)) - and isinstance(p.type.inner_type, ast_nodes.ValueType) - and p.type.inner_type.name == "char" - ): - const_qualifier = "const " if p.type.inner_type.is_const else "" - native_params.append( - f"{p.name}.as<{const_qualifier}std::string>().data()" - ) - else: - raise TypeError( - f"Unhandled parameter type for conversion: {p.type} for param" - f" '{p.name}'" - ) - return native_params - - -def get_compatible_return_call( - func: ast_nodes.FunctionDecl, invoker: str -) -> str: - """Generates embind compatible return value conversion.""" - - if return_is_value_of_type(func, {"void"}): - return invoker - if isinstance(func.return_type, ast_nodes.PointerType) and isinstance( - func.return_type.inner_type, ast_nodes.ValueType - ): - if func.return_type.inner_type.name == "char": - return f"return std::string({invoker})" - elif func.return_type.inner_type.name == "mjString": - return f"return *{invoker}" - if return_is_pointer_to_struct(func): - return get_converted_struct_to_class(func, invoker) - if return_is_value_of_type(func, PRIMITIVE_TYPES): - return f"return {invoker}" - raise RuntimeError( - "Failed to calculate return value conversion for function" - f" {func.name} that returns '{func.return_type}'" - ) - - -def get_compatible_return_type(func: ast_nodes.FunctionDecl) -> str: - """Creates embind compatible return type.""" - - if ( - isinstance(func.return_type, ast_nodes.PointerType) - and isinstance(func.return_type.inner_type, ast_nodes.ValueType) - and func.return_type.inner_type.name in ["char", "mjString"] - ): - return "std::string" - if ( - isinstance(func.return_type, ast_nodes.PointerType) - and isinstance(func.return_type.inner_type, ast_nodes.ValueType) - and func.return_type.inner_type.name not in PRIMITIVE_TYPES - ): - const_qualifier = get_const_qualifier(func) - return f"""{const_qualifier}std::optional<{uppercase_first_letter(func.return_type.inner_type.name)}>""" - if ( - isinstance(func.return_type, ast_nodes.ValueType) - and func.return_type.name in PRIMITIVE_TYPES - ): - return f"{func.return_type.name}" - return "val" - - -def get_converted_struct_to_class( - func: ast_nodes.FunctionDecl, invoker: str -) -> str: - """Generates a C++ function invocation for a struct return-type function.""" - - const_qualifier = get_const_qualifier(func) - return_type = cast(ast_nodes.PointerType, func.return_type) - struct_name = cast(ast_nodes.ValueType, return_type.inner_type).name - class_constructor = uppercase_first_letter(struct_name) - return_str = f"{class_constructor}(result)" - return f"""{const_qualifier}{struct_name}* result = {invoker}; - if (result == nullptr) {{ - return std::nullopt; - }} - return {return_str}""" - - -def is_excluded_function_name(func_name: str) -> bool: - """Checks if a function name should be excluded from direct binding.""" - return ( - func_name.startswith("mjr_") - or func_name.startswith("mjui_") - or func_name in constants.SKIPPED_FUNCTIONS - ) diff --git a/wasm/codegen/helpers/structs.py b/wasm/codegen/helpers/structs.py deleted file mode 100644 index 810e6b89..00000000 --- a/wasm/codegen/helpers/structs.py +++ /dev/null @@ -1,714 +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 -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 - - -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 = "" - - # Initialization code for fields that require it - initialization: str = "" - - # Statement to reset the inner pointer when copying the field - ptr_copy_reset: str = "" - - # Whether the field is a primitive or fixed size - is_primitive_or_fixed_size: bool = False - - # 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 = "" - - -@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 - - # Struct bindings code - bindings: str = "" - - -def _simple_property_binding( - field: ast_nodes.StructFieldDecl, - struct_wrapper_name: str, - setter: bool = False, - reference: bool = False, -) -> str: - """Builds the C++ code for a simple property binding.""" - 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})' - - -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) - - 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( - 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, - ) - - 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, - ) - - 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, - ) - - elif isinstance(f.type, ast_nodes.PointerType): - - inner_type_name = ( - f.type.inner_type.name - if isinstance(f.type.inner_type, ast_nodes.ValueType) - else "" - ) - ptr_field_expr = f"ptr_->{f.name}" - array_size_str = "" - - 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[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=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=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() - ): - ptr_field = cast(ast_nodes.PointerType, f.type) - wrapper_field_name = common.uppercase_first_letter( - cast(ast_nodes.ValueType, ptr_field.inner_type).name - ) - return WrappedFieldData( - 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})", - ) - - 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 WrappedFieldData( - definition=builder.to_string(), - typename=_get_field_struct_type(f.type), - binding=_simple_property_binding(f, w), - ) - - # SHOULD NOT OCCUR - print("Error: field {f.name} not properly handled") - return WrappedFieldData( - definition=f"// Error: field {f.name} not properly handled.", - typename=_get_field_struct_type(f.type), - binding=f"// Error: field {f.name} not properly handled.", - ) - - -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: - member_type = cast(ast_nodes.StructFieldDecl, field).type - if isinstance(member_type, (ast_nodes.ArrayType, ast_nodes.PointerType)): - member_type = member_type.inner_type - if isinstance(member_type, ast_nodes.ValueType): - return member_type.name in constants.STRUCTS_TO_BIND - 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.""" - s = struct_name - w = common.uppercase_first_letter(s) - - shallow_copy = use_shallow_copy(wrapped_fields) - - builder = code_builder.CodeBuilder() - with builder.struct(f"{w}"): - builder.line(f"explicit {w}({s} *ptr);") - builder.line(f"~{w}();") - - if not is_mjs: - builder.line(f"{w}();") - - if shallow_copy and not is_mjs: - builder.line(f"{w}(const {w} &);") - builder.line(f"{w} &operator=(const {w} &);") - builder.line(f"std::unique_ptr<{w}> copy();") - - builder.line(f"{s}* get() const;") - builder.line(f"void set({s}* ptr);") - - 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.private() - builder.line(f"{s}* ptr_;") - if not is_mjs: - builder.line("bool owned_ = false;") - - if is_mjs and fields_with_init: - builder.public() - for field in fields_with_init: - if field.definition: - builder.line(f"{field.definition}") - - return builder.to_string() + ";" - - -def _default_function_statement(struct_name: str) -> str: - """Returns the default function name for the given struct.""" - if struct_name == "mjvGeom": - f = "mjv_initGeom(ptr_, mjGEOM_NONE, nullptr, nullptr, nullptr, nullptr);" - return f - elif struct_name in constants.ANONYMOUS_STRUCTS.keys(): - return "" - elif struct_name in constants.NO_DEFAULT_CONSTRUCTORS: - return "" - elif struct_name.startswith("mjs"): - return f"mjs_default{struct_name.removeprefix('mjs')}(ptr_);" - elif struct_name.startswith("mjv"): - return f"mjv_default{struct_name.removeprefix('mjv')}(ptr_);" - elif struct_name.startswith("mj"): - return f"mj_default{struct_name.removeprefix('mj')}(ptr_);" - - return "" - - -def _delete_ptr_statement(struct_name: str) -> str: - """Returns the delete function name for the given struct.""" - if struct_name == "mjVFS": - return "mj_deleteVFS(ptr_);" - - return "delete ptr_;" - - -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, - ) - - is_anonymous_struct = struct_name in constants.ANONYMOUS_STRUCTS.keys() - is_hardcoded_wrapper_struct = ( - common.uppercase_first_letter(struct_name) in constants.MANUAL_STRUCTS - ) - - if ( - not is_hardcoded_wrapper_struct - and struct_info - and not _has_nested_wrapper_members(struct_info) - or is_anonymous_struct - ): - 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.""" - # These structs require specific function calls for creation and/or deletion - # which, for now, are hardcoded in the template file. - if struct_name in [ - "mjData", - "mjModel", - "mjvScene", - "mjSpec", - ]: - return "" - - s = struct_name - w = common.uppercase_first_letter(s) - is_mjs = w.startswith("Mjs") - - fields_with_init = _find_fields_with_init(wrapped_fields) - shallow_copy = use_shallow_copy(wrapped_fields) - - fields_init = "" - if fields_with_init: - fields_init = "".join( - field_with_init.initialization for field_with_init in fields_with_init - ) - - builder = code_builder.CodeBuilder() - - # constructor passing native ptr - with builder.function(f"{w}::{w}({s} *ptr) : ptr_(ptr){fields_init}"): - pass - - # destructor - with builder.function(f"{w}::~{w}()"): - if not is_mjs: - with builder.block("if (owned_ && ptr_)"): - delete_ptr = _delete_ptr_statement(s) - builder.line(delete_ptr) - - if not is_mjs: - # default constructor - with builder.function(f"{w}::{w}() : ptr_(new {s}){fields_init}"): - builder.line("owned_ = true;") - default_func = _default_function_statement(s) - if default_func: - builder.line(default_func) - - if shallow_copy and not is_mjs: - # copy constructor - with builder.function(f"{w}::{w}(const {w} &other) : {w}()"): - builder.line("*ptr_ = *other.get();") - 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.function(f"{w}& {w}::operator=(const {w} &other)"): - with builder.block("if (this == &other)"): - builder.line("return *this;") - builder.line("*ptr_ = *other.get();") - 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;") - - # explicit copy function - with builder.function(f"std::unique_ptr<{w}> {w}::copy()"): - builder.line(f"return std::make_unique<{w}>(*this);") - - # C struct getter/setter - with builder.function(f"{s}* {w}::get() const"): - builder.line("return ptr_;") - with builder.function(f"void {w}::set({s}* ptr)"): - builder.line("ptr_ = ptr;") - - return builder.to_string() - - -def _build_struct_bindings( - struct_name: str, - wrapped_fields: List[WrappedFieldData], -): - """Builds the C++ bindings for a struct.""" - w = common.uppercase_first_letter(struct_name) - is_mjs = w.startswith("Mjs") - - builder = code_builder.CodeBuilder() - with builder.block( - header_line=f'emscripten::class_<{w}>("{w}")', braces=False - ): - if w == "MjData": - builder.line(".constructor()") - builder.line(".constructor()") - elif w == "MjModel": - builder.line( - '.class_function("loadFromXML", &loadFromXML, take_ownership())' - ) - builder.line(".constructor()") - elif w == "MjSpec": - builder.line(".constructor()") - elif w == "MjvScene": - builder.line(".constructor()") - builder.line(".constructor<>()") - elif not is_mjs: - builder.line(".constructor<>()") - - shallow_copy = use_shallow_copy(wrapped_fields) - if shallow_copy and not is_mjs: - builder.line(f'.function("copy", &{w}::copy, take_ownership())') - - for field in wrapped_fields[:-1]: - if field.binding: - builder.line(field.binding) - if wrapped_fields: - builder.line(f"{wrapped_fields[-1].binding};") - - 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: - s = struct_name - w = common.uppercase_first_letter(s) - - if s in introspect_structs: - struct_fields = introspect_structs[s].fields - elif s in constants.ANONYMOUS_STRUCTS: - anonymous_struct = _get_anonymous_struct_field(s) - if not anonymous_struct or not isinstance( - anonymous_struct.type, ast_nodes.AnonymousStructDecl - ): - raise RuntimeError(f"Anonymous struct not found: {s}") - struct_fields = anonymous_struct.type.fields - else: - raise RuntimeError(f"Struct not found: {s}") - - wrapped_fields: List[WrappedFieldData] = [] - for field in struct_fields: - wrapped_fields.append(_generate_field_data(field, w)) - - wrap_data = WrappedStructData( - wrap_name=w, - wrapped_fields=wrapped_fields, - wrapped_header=build_struct_header(s, wrapped_fields), - wrapped_source=build_struct_source(s, wrapped_fields), - bindings=_build_struct_bindings(s, wrapped_fields), - ) - - wrapped_structs[s] = 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_wrappers: dict[str, WrappedStructData], -) -> List[str]: - """Sorts structs based on their field dependencies using topological sort. - - Structs with no dependencies on other structs in the list come first. - Struct A has a dependency on struct B if struct A has a field where the - underlying_type is B. Note that this definition is stricter than the C++ - struct dependency criterion where forward declarations can be used to - eliminate dependencies A and B if A only has a pointer to B. - - Args: - struct_wrappers: A dictionary mapping struct names to their - WrappedStructData. - - 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_names = struct_wrappers.keys() - struct_set = set(struct_names) - sorted_struct_names = sorted(struct_names) - - for struct_name in sorted_struct_names: - for field in struct_wrappers[struct_name].wrapped_fields: - - field_type_name = field.typename - 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/bindings_diff_test.py b/wasm/codegen/tests/bindings_diff_test.py similarity index 91% rename from wasm/codegen/bindings_diff_test.py rename to wasm/codegen/tests/bindings_diff_test.py index cf77f017..2281778a 100644 --- a/wasm/codegen/bindings_diff_test.py +++ b/wasm/codegen/tests/bindings_diff_test.py @@ -16,7 +16,7 @@ from pathlib import Path from absl.testing import absltest -from wasm.codegen import update +from wasm.codegen.generators import binding_builder ERROR_MESSAGE = """ The file '{}' needs to be updated, please run: @@ -31,7 +31,7 @@ class BindingsDiffTest(absltest.TestCase): self.generated_src = f.read() self.template_path_cc = SCRIPT_DIR / 'templates/bindings.cc' - self.builder = update.BindingBuilder(self.template_path_cc) + self.builder = binding_builder.BindingBuilder(self.template_path_cc) generator_output = ( self.builder.set_enums().set_structs().set_functions().to_string() diff --git a/wasm/codegen/coverage_test.py b/wasm/codegen/tests/coverage_test.py similarity index 97% rename from wasm/codegen/coverage_test.py rename to wasm/codegen/tests/coverage_test.py index fcbb8559..fbf5c15e 100644 --- a/wasm/codegen/coverage_test.py +++ b/wasm/codegen/tests/coverage_test.py @@ -33,9 +33,9 @@ from absl.testing import absltest from introspect import functions as introspect_functions from introspect import structs as introspect_structs -from wasm.codegen.helpers import common -from wasm.codegen.helpers import constants -from wasm.codegen.helpers import functions +from wasm.codegen.generators import common +from wasm.codegen.generators import constants +from wasm.codegen.generators import functions def _get_resource_content(file_path: str) -> str: diff --git a/wasm/codegen/enums_test_generator.py b/wasm/codegen/tests/enums_test_generator.py similarity index 97% rename from wasm/codegen/enums_test_generator.py rename to wasm/codegen/tests/enums_test_generator.py index dd3a69a5..f380aabc 100644 --- a/wasm/codegen/enums_test_generator.py +++ b/wasm/codegen/tests/enums_test_generator.py @@ -18,7 +18,7 @@ import textwrap from introspect import enums as introspect_enums -from wasm.codegen.helpers import common +from wasm.codegen.generators import common def generate_typescript_enum_tests(): diff --git a/wasm/codegen/helpers/helpers_test.py b/wasm/codegen/tests/generators_test.py similarity index 90% rename from wasm/codegen/helpers/helpers_test.py rename to wasm/codegen/tests/generators_test.py index a539f854..640125ec 100644 --- a/wasm/codegen/helpers/helpers_test.py +++ b/wasm/codegen/tests/generators_test.py @@ -17,11 +17,12 @@ from absl.testing import absltest from introspect import ast_nodes -from wasm.codegen.helpers import code_builder -from wasm.codegen.helpers import common -from wasm.codegen.helpers import constants -from wasm.codegen.helpers import functions -from wasm.codegen.helpers import structs +from wasm.codegen.generators import code_builder +from wasm.codegen.generators import common +from wasm.codegen.generators import constants +from wasm.codegen.generators import enums +from wasm.codegen.generators import functions +from wasm.codegen.generators import structs class CodeBuilderTest(absltest.TestCase): @@ -877,5 +878,88 @@ emscripten::val multi_dim_array() const { ) +class EnumsGeneratorTest(absltest.TestCase): + + def test_generate_enum_bindings(self): + + generator = enums.Generator({ + "TestEnum": ast_nodes.EnumDecl( + name="TestEnum", + declname="enum TestEnum_", + values={"FIRST_VAL": 0, "SECOND_VAL": 1, "THIRD_VAL": 2}, + ), + "AnotherEnum": ast_nodes.EnumDecl( + name="AnotherEnum", + declname="enum AnotherEnum_", + values={"ALPHA": 100, "BETA": 200}, + ), + "EmptyEnum": ast_nodes.EnumDecl( + name="EmptyEnum", + declname="enum EmptyEnum_", + values={}, + ), + }) + + expected_code = """ enum_("TestEnum") + .value("FIRST_VAL", FIRST_VAL) + .value("SECOND_VAL", SECOND_VAL) + .value("THIRD_VAL", THIRD_VAL); + + enum_("AnotherEnum") + .value("ALPHA", ALPHA) + .value("BETA", BETA); + + enum_("EmptyEnum");""" + + markers_and_content = generator.generate() + actual_code = "\n\n".join(markers_and_content[0][1]) + + self.assertEqual(actual_code, expected_code) + + +class FunctionsGeneratorTest(absltest.TestCase): + + def setUp(self): + super().setUp() + self.generator = functions.Generator({}) + self.int_type = ast_nodes.ValueType(name="int") + + def test_generate_function_binding_simple_case(self): + func_simple_void = ast_nodes.FunctionDecl( + name="do_nothing", + return_type=ast_nodes.ValueType(name="void"), + parameters=tuple(), + doc="doc", + ) + self.assertEqual( + self.generator._generate_function_binding(func_simple_void), + 'function("do_nothing", &do_nothing);\n', + ) + + def test_generate_direct_bindable_functions_simple_filter(self): + direct_bind = ast_nodes.FunctionDecl( + name="direct_bind", + return_type=self.int_type, + parameters=( + ast_nodes.FunctionParameterDecl(name="val", type=self.int_type), + ), + doc="doc", + ) + needs_wrap = ast_nodes.FunctionDecl( + name="needs_wrap", + return_type=ast_nodes.PointerType(inner_type=self.int_type), + parameters=tuple(), + doc="doc", + ) + self.generator = functions.Generator({ + "direct1": direct_bind, + "wrapped1": needs_wrap, + }) + + generated_code = self.generator._generate_direct_bindable_functions() + self.assertIn('function("direct_bind", &direct_bind);\n', generated_code) + self.assertNotIn('function("needs_wrap", &needs_wrap);\n', generated_code) + + if __name__ == "__main__": absltest.main() diff --git a/wasm/codegen/update.py b/wasm/codegen/update.py index 776e2d2e..f6741a61 100644 --- a/wasm/codegen/update.py +++ b/wasm/codegen/update.py @@ -14,73 +14,14 @@ """Generates Javascript/TypeScript bindings for MuJoCo.""" -import os - -from introspect import ast_nodes -from introspect import enums as introspect_enums -from introspect import functions as introspect_functions - -from wasm.codegen.generators import enums -from wasm.codegen.generators import functions -from wasm.codegen.generators import structs -from wasm.codegen.helpers import common -from wasm.codegen.helpers import constants as _constants -from wasm.codegen.helpers import functions as function_utils - - -class BindingBuilder: - """Builds WASM bindings for MuJoCo.""" - - def __init__( - self, - template_path_cc: str, - ): - with open(template_path_cc, "r") as f: - self.content_cc = f.readlines() - - self.markers_and_content = [] - - def set_enums(self): - """Generates and sets the enum bindings.""" - generator = enums.Generator(introspect_enums.ENUMS) - self.markers_and_content += generator.generate() - return self - - def set_structs(self): - """Generates and sets the struct bindings.""" - generator = structs.Generator() - self.markers_and_content += generator.generate() - return self - - def set_functions(self): - """Generates and sets the function wrappers and bindings.""" - - functions_to_bind: dict[str, ast_nodes.FunctionDecl] = {} - for name, func in introspect_functions.FUNCTIONS.items(): - if not function_utils.is_excluded_function_name(name): - if name not in _constants.BOUNDCHECK_FUNCS: - functions_to_bind[name] = func - - generator = functions.Generator(functions_to_bind) - self.markers_and_content += generator.generate() - return self - - def to_string(self) -> str: - for marker, content in self.markers_and_content: - self.content_cc = common.replace_lines_containing_marker( - self.content_cc, marker, content - ) - return "".join(self.content_cc) - - def build(self, generated_path_cc: str): - common.write_to_file(generated_path_cc, self.to_string()) +from wasm.codegen.generators import binding_builder if __name__ == "__main__": template_file = "wasm/codegen/templates/bindings.cc" generated_file = "wasm/codegen/generated/bindings.cc" - builder = BindingBuilder(template_file) + builder = binding_builder.BindingBuilder(template_file) builder.set_enums() builder.set_structs() builder.set_functions() diff --git a/wasm/tests/benchmark_test.ts b/wasm/tests/benchmark_test.ts index 9928ff55..3fac15d1 100644 --- a/wasm/tests/benchmark_test.ts +++ b/wasm/tests/benchmark_test.ts @@ -59,7 +59,7 @@ describe('MuJoCo WASM Benchmark Tests', () => { } // Warmup JIT compiler to get stable results - for (let i = 0; i < 1000; i++) { + for (let i = 0; i < 100; i++) { const sortedState = func(state); verifySorted(sortedState, kind); }