diff --git a/wasm/codegen/binding_builder.py b/wasm/codegen/binding_builder.py index 78bcb130..1f61766a 100644 --- a/wasm/codegen/binding_builder.py +++ b/wasm/codegen/binding_builder.py @@ -14,13 +14,13 @@ """Builds WASM bindings for MuJoCo.""" +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 @@ -32,73 +32,45 @@ class BindingBuilder: def __init__( self, template_path_cc: str, - generated_path_cc: str, ): - self.generated_path_cc = generated_path_cc with open(template_path_cc, "r") as f: self.content_cc = f.readlines() - filtered_functions = { - name: func - for name, func in introspect_functions.FUNCTIONS.items() - if not function_utils.is_excluded_function_name(name) - and name not in _constants.BOUNDCHECK_FUNCS - } - self.enums_generator = enums.Generator(introspect_enums.ENUMS) - self.functions_generator = functions.Generator(filtered_functions) - self.structs_generator = structs.Generator() + self.markers_and_content = [] def set_enums(self): """Generates and sets the enum bindings.""" - enum_bindings = self.enums_generator.generate() - self.content_cc = common.replace_lines_containing_marker( - self.content_cc, - "// {{ ENUM_BINDINGS }}", - 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.""" - - # Generate struct header bindings. - struct_hdr_markers_and_content = self.structs_generator.generate_header() - for marker, content in struct_hdr_markers_and_content: - self.content_cc = common.replace_lines_containing_marker( - self.content_cc, marker, content - ) - - # Generate struct source bindings. - struct_src_markers_and_content = ( - self.structs_generator.generate_source() - ) - for marker, content in struct_src_markers_and_content: - self.content_cc = common.replace_lines_containing_marker( - self.content_cc, marker, content - ) - + generator = structs.Generator() + self.markers_and_content += generator.generate() return self def set_functions(self): """Generates and sets the function wrappers and bindings.""" - wrapper_functions, function_bindings = ( - self.functions_generator.generate() - ) - self.content_cc = common.replace_lines_containing_marker( - self.content_cc, - "// {{ WRAPPER_FUNCTIONS }}", - wrapper_functions, - ) - self.content_cc = common.replace_lines_containing_marker( - self.content_cc, - "// {{ FUNCTION_BINDINGS }}", - function_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 build(self): - """Writes the generated content to the output files.""" - common.write_to_file(self.generated_path_cc, "".join(self.content_cc)) - - def to_string_source(self) -> str: + 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): + """Writes the generated content to the output files.""" + + common.write_to_file(generated_path_cc, self.to_string()) diff --git a/wasm/codegen/bindings_diff_test.py b/wasm/codegen/bindings_diff_test.py index 4a0cc820..edbfd732 100644 --- a/wasm/codegen/bindings_diff_test.py +++ b/wasm/codegen/bindings_diff_test.py @@ -29,21 +29,16 @@ class BindingsDiffTest(absltest.TestCase): super().setUp() SCRIPT_DIR = Path(__file__).parent - with open(SCRIPT_DIR / 'generated/bindings.h', 'r') as f: - self.generated_hdr = f.read() with open(SCRIPT_DIR / 'generated/bindings.cc', 'r') as f: self.generated_src = f.read() self.template_path_cc = SCRIPT_DIR / 'templates/bindings.cc' - self.generated_path_cc = SCRIPT_DIR / 'generated/bindings.cc' - self.builder = binding_builder.BindingBuilder( - self.template_path_cc, - self.generated_path_cc, - ) + self.builder = binding_builder.BindingBuilder(self.template_path_cc) def test_bindings_source(self): - generator_output = (self.builder.set_enums().set_structs().set_functions(). - to_string_source()) + generator_output = ( + self.builder.set_enums().set_structs().set_functions().to_string() + ) self.assertEqual( generator_output, self.generated_src, diff --git a/wasm/codegen/generated/bindings.cc b/wasm/codegen/generated/bindings.cc index 25b22757..451890dc 100644 --- a/wasm/codegen/generated/bindings.cc +++ b/wasm/codegen/generated/bindings.cc @@ -11846,6 +11846,7 @@ std::optional mjs_asPlugin_wrapper(MjsElement& element) { return MjsPlugin(result); } + void mju_printMatSparse_wrapper(const NumberArray& mat, const NumberArray& rownnz, const NumberArray& rowadr, const NumberArray& colind) { UNPACK_ARRAY(mjtNum, mat); @@ -12877,6 +12878,7 @@ EMSCRIPTEN_BINDINGS(mujoco_functions) { function("mjs_asTexture", &mjs_asTexture_wrapper); function("mjs_asMaterial", &mjs_asMaterial_wrapper); function("mjs_asPlugin", &mjs_asPlugin_wrapper); + function("error", &error_wrapper); function("mju_printMatSparse", &mju_printMatSparse_wrapper); function("mj_solveM", &mj_solveM_wrapper); diff --git a/wasm/codegen/generators/enums.py b/wasm/codegen/generators/enums.py index ecbb66d0..b9f3a7a4 100644 --- a/wasm/codegen/generators/enums.py +++ b/wasm/codegen/generators/enums.py @@ -14,7 +14,7 @@ """Generates Embind bindings for MuJoCo enums.""" -from typing import Mapping +from typing import Mapping, Optional from introspect import ast_nodes @@ -38,10 +38,13 @@ class Generator: code += ";" return code - def generate(self) -> str: + def generate(self) -> list[tuple[str, list[str]]]: """Generates all Embind code for the provided enums.""" code = [] for enum in self.enums.values(): code.append(self._generate_enum_binding(enum)) - return "\n\n".join(code) + "\n" + + content = "\n\n".join(code) + marker = "// {{ ENUM_BINDINGS }}" + return [(marker, [content])] diff --git a/wasm/codegen/generators/functions.py b/wasm/codegen/generators/functions.py index 1ef2e52e..5e7d8f8f 100644 --- a/wasm/codegen/generators/functions.py +++ b/wasm/codegen/generators/functions.py @@ -14,7 +14,7 @@ """Generates Embind bindings for MuJoCo functions.""" -from typing import List, Mapping +from typing import List, Mapping, Optional from introspect import ast_nodes @@ -76,11 +76,14 @@ class Generator: return result - def generate(self) -> tuple[str, str]: + def generate(self) -> list[tuple[str, list[str]]]: """Generates the bindings file for all functions.""" wrapper_functions = self._generate_wrappers() function_bindings = self._generate_direct_bindable_functions() function_bindings += self._generate_wrapper_bindable_functions() - return wrapper_functions, function_bindings + return [ + ("// {{ WRAPPER_FUNCTIONS }}", [wrapper_functions]), + ("// {{ FUNCTION_BINDINGS }}", [function_bindings]), + ] diff --git a/wasm/codegen/generators/generators_test.py b/wasm/codegen/generators/generators_test.py index 8ef097a0..bf61e90f 100644 --- a/wasm/codegen/generators/generators_test.py +++ b/wasm/codegen/generators/generators_test.py @@ -13,7 +13,6 @@ # limitations under the License. from absl.testing import absltest - from introspect import ast_nodes from wasm.codegen.generators import enums @@ -51,13 +50,14 @@ class EnumsGeneratorTest(absltest.TestCase): .value("ALPHA", ALPHA) .value("BETA", BETA); - enum_("EmptyEnum"); -""" + enum_("EmptyEnum");""" - actual_code = generator.generate() + 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): diff --git a/wasm/codegen/generators/structs.py b/wasm/codegen/generators/structs.py index bea5b060..2d969064 100644 --- a/wasm/codegen/generators/structs.py +++ b/wasm/codegen/generators/structs.py @@ -23,17 +23,15 @@ from wasm.codegen.helpers import structs class Generator: """Generates C++ code for binding and wrapping MuJoCo structs.""" - def __init__(self): + def generate(self) -> list[tuple[str, list[str]]]: + """Generates C++ header file for binding and wrapping MuJoCo structs.""" + # 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( constants.STRUCTS_TO_BIND ) - def generate_header( - self - ) -> list[tuple[str, list[Optional[str]]]]: - """Generates C++ header file for binding and wrapping MuJoCo structs.""" autogenned_struct_definitions = [] markers_and_content = [] @@ -60,22 +58,18 @@ class Generator: "// {{ AUTOGENNED_STRUCT_DEFINITIONS }}", autogenned_struct_definitions, )) - return markers_and_content - def generate_source(self) -> list[tuple[str, list[str]]]: - """Generates C++ source file for binding and wrapping MuJoCo structs.""" - constructors = [ - ( - f"// INSERT-GENERATED-{struct_data.wrap_name}-CONSTRUCTOR", - [struct_data.wrapped_source], - ) - for _, struct_data in self.structs_to_bind_data.items() - ] - properties = [ - ( - f"// INSERT-GENERATED-{struct_data.wrap_name}-BINDINGS", - [l.binding for l in struct_data.wrapped_fields], - ) - for _, struct_data in self.structs_to_bind_data.items() - ] - return constructors + properties + for _, struct_data in self.structs_to_bind_data.items(): + # Bindings + markers_and_content.append(( + f"// INSERT-GENERATED-{struct_data.wrap_name}-BINDINGS", + [l.binding for l in struct_data.wrapped_fields], + )) + + # Special member functions + markers_and_content.append(( + f"// INSERT-GENERATED-{struct_data.wrap_name}-CONSTRUCTOR", + [struct_data.wrapped_source], + )) + + return markers_and_content diff --git a/wasm/codegen/helpers/constants.py b/wasm/codegen/helpers/constants.py index 1839dc4f..10e9ccb8 100644 --- a/wasm/codegen/helpers/constants.py +++ b/wasm/codegen/helpers/constants.py @@ -14,7 +14,7 @@ """Constants used in the code generation process.""" -from typing import List, Set, Dict +from typing import Dict, List, Set from introspect import structs as introspect_structs PRIMITIVE_TYPES: Set[str] = { @@ -189,16 +189,16 @@ _UTILITY_FUNCTIONS: List[str] = [ # List of functions that should be skipped during the code generation process. SKIPPED_FUNCTIONS: List[str] = ( - _CLASS_METHODS + - _THREAD_FUNCTIONS + - _MEMORY_FUNCTIONS + - _PLUGIN_FUNCTIONS + - _GETTERS_AND_SETTERS + - _VISUAL_FUNCTIONS + - _ASSET_CACHE_FUNCTIONS + - _VFS_FUNCTIONS + - _WRITABLE_ERROR + - _UTILITY_FUNCTIONS + _CLASS_METHODS + + _THREAD_FUNCTIONS + + _MEMORY_FUNCTIONS + + _PLUGIN_FUNCTIONS + + _GETTERS_AND_SETTERS + + _VISUAL_FUNCTIONS + + _ASSET_CACHE_FUNCTIONS + + _VFS_FUNCTIONS + + _WRITABLE_ERROR + + _UTILITY_FUNCTIONS ) # Functions that require special wrappers to infer sizes and make additional diff --git a/wasm/codegen/helpers/structs.py b/wasm/codegen/helpers/structs.py index 76fc83a9..68b58e42 100644 --- a/wasm/codegen/helpers/structs.py +++ b/wasm/codegen/helpers/structs.py @@ -635,9 +635,9 @@ def build_struct_source( 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;") diff --git a/wasm/codegen/update.py b/wasm/codegen/update.py index e75f7786..b32684a3 100644 --- a/wasm/codegen/update.py +++ b/wasm/codegen/update.py @@ -28,8 +28,8 @@ def generate_all_bindings(): template_path_cc, generated_path_cc = common.get_file_path( "templates", "generated", "bindings.cc" ) - builder = binding_builder.BindingBuilder(template_path_cc, generated_path_cc) - builder.set_enums().set_structs().set_functions().build() + builder = binding_builder.BindingBuilder(template_path_cc) + builder.set_enums().set_structs().set_functions().build(generated_path_cc) if __name__ == "__main__":