From 217317ccbaed497c3a19ba5f6706b02c0e179e0a Mon Sep 17 00:00:00 2001 From: Google DeepMind Date: Mon, 24 Nov 2025 09:35:06 -0800 Subject: [PATCH] Change WASM generators from classes to functions. PiperOrigin-RevId: 836257289 Change-Id: I7a3c1d7bb545353a3d278330582e33c0551792cb --- wasm/codegen/generators/binding_builder.py | 10 +- wasm/codegen/generators/enums.py | 38 ++++---- wasm/codegen/generators/functions.py | 44 ++++----- wasm/codegen/generators/structs.py | 107 ++++++++++----------- wasm/codegen/tests/generators_test.py | 31 +++--- 5 files changed, 105 insertions(+), 125 deletions(-) diff --git a/wasm/codegen/generators/binding_builder.py b/wasm/codegen/generators/binding_builder.py index 9ff46b1c..cd9ebff7 100644 --- a/wasm/codegen/generators/binding_builder.py +++ b/wasm/codegen/generators/binding_builder.py @@ -17,6 +17,7 @@ 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 @@ -36,14 +37,12 @@ class BindingBuilder: def set_enums(self): """Generates and sets the enum bindings.""" - generator = enums.Generator(introspect_enums.ENUMS) - self.markers_and_content += generator.generate() + self.markers_and_content += enums.generate(introspect_enums.ENUMS) return self def set_structs(self): """Generates and sets the struct bindings.""" - generator = structs.Generator() - self.markers_and_content += generator.generate() + self.markers_and_content += structs.generate(constants.STRUCTS_TO_BIND) return self def set_functions(self): @@ -54,8 +53,7 @@ class BindingBuilder: if not functions.is_excluded_function_name(name): functions_to_bind[name] = func - generator = functions.Generator(functions_to_bind) - self.markers_and_content += generator.generate() + self.markers_and_content += functions.generate(functions_to_bind) return self def to_string(self) -> str: diff --git a/wasm/codegen/generators/enums.py b/wasm/codegen/generators/enums.py index 398f070b..f2a1b5fa 100644 --- a/wasm/codegen/generators/enums.py +++ b/wasm/codegen/generators/enums.py @@ -21,26 +21,22 @@ from introspect import ast_nodes from wasm.codegen.generators import code_builder -class Generator: - """Generates Embind code for MuJoCo enums.""" +def generate( + enums: Mapping[str, ast_nodes.EnumDecl], +) -> list[tuple[str, list[str]]]: + """Generates all Embind code for the provided enums.""" - def __init__(self, enums: Mapping[str, ast_nodes.EnumDecl]): - self.enums = enums + builder = code_builder.CodeBuilder() + with builder.block('EMSCRIPTEN_BINDINGS(mujoco_enums)'): + for e in enums.values(): + if e.values: # Skip empty enums. + with builder.block(f'enum_<{e.name}>("{e.name}")', braces=False): + names = list(e.values.keys()) + for name in names[:-1]: + builder.line(f'.value("{name}", {name})') + builder.line(f'.value("{names[-1]}", {names[-1]});') + builder.newline() - def generate(self) -> list[tuple[str, list[str]]]: - """Generates all Embind code for the provided enums.""" - - builder = code_builder.CodeBuilder() - with builder.block('EMSCRIPTEN_BINDINGS(mujoco_enums)'): - for e in self.enums.values(): - if e.values: # Skip empty enums. - with builder.block(f'enum_<{e.name}>("{e.name}")', braces=False): - names = list(e.values.keys()) - for name in names[:-1]: - builder.line(f'.value("{name}", {name})') - builder.line(f'.value("{names[-1]}", {names[-1]});') - builder.newline() - - content = builder.to_string() - marker = '// {{ ENUM_BINDINGS }}' - return [(marker, [content])] + content = builder.to_string() + marker = '// {{ ENUM_BINDINGS }}' + return [(marker, [content])] diff --git a/wasm/codegen/generators/functions.py b/wasm/codegen/generators/functions.py index 0bc69ee3..d587202c 100644 --- a/wasm/codegen/generators/functions.py +++ b/wasm/codegen/generators/functions.py @@ -14,7 +14,7 @@ """Helper functions for processing and generating bindings for MuJoCo functions.""" -from typing import List, Mapping, Set, Tuple, cast +from typing import Mapping, Tuple, cast from introspect import ast_nodes @@ -288,30 +288,24 @@ def is_excluded_function_name(func_name: str) -> bool: ) -class Generator: +def generate( + functions: Mapping[str, ast_nodes.FunctionDecl], +) -> list[tuple[str, list[str]]]: """Generates Embind bindings for MuJoCo functions.""" + wrapper_functions = [] + for func in functions.values(): + if should_be_wrapped(func): + if func.name not in constants.MANUAL_WRAPPER_FUNCTIONS: + wrapper_functions.append(generate_function_wrapper(func)) + wrapper_content = "\n\n".join(wrapper_functions) - def __init__(self, functions: Mapping[str, ast_nodes.FunctionDecl]): - self.functions = functions + function_bindings = [] + for func in functions.values(): + suffix = "_wrapper" if should_be_wrapped(func) else "" + function_bindings.append(f'function("{func.name}", &{func.name}{suffix});') + bindings_content = "\n".join(function_bindings) - def generate(self) -> list[tuple[str, list[str]]]: - """Generates the bindings file for all functions.""" - - function_wrappers = [] - for func in self.functions.values(): - if should_be_wrapped(func): - if func.name not in constants.MANUAL_WRAPPER_FUNCTIONS: - function_wrappers.append(generate_function_wrapper(func)) - - function_bindings = [] - for func in self.functions.values(): - js_name = func.name - cpp_func = func.name - if should_be_wrapped(func): - cpp_func += "_wrapper" - function_bindings.append(f'function("{js_name}", &{cpp_func});') - - return [ - ("// {{ WRAPPER_FUNCTIONS }}", ["\n\n".join(function_wrappers)]), - ("// {{ FUNCTION_BINDINGS }}", ["\n".join(function_bindings)]), - ] + return [ + ("// {{ WRAPPER_FUNCTIONS }}", [wrapper_content]), + ("// {{ FUNCTION_BINDINGS }}", [bindings_content]), + ] diff --git a/wasm/codegen/generators/structs.py b/wasm/codegen/generators/structs.py index 1a784b39..5f6e2ac5 100644 --- a/wasm/codegen/generators/structs.py +++ b/wasm/codegen/generators/structs.py @@ -719,66 +719,61 @@ def sort_structs_by_dependency( ) -class Generator: - """Generates C++ code for binding and wrapping MuJoCo structs.""" +def generate(struct_to_bind: List[str]) -> list[tuple[str, list[str]]]: + """Generates C++ header file for binding and wrapping MuJoCo structs.""" - 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 + structs_to_bind_data = generate_wasm_bindings(struct_to_bind) - # Traverse the introspect dictionary to get the field - # wrapper/bindings statements set up for each struct - self.structs_to_bind_data = generate_wasm_bindings( - constants.STRUCTS_TO_BIND + autogenned_struct_definitions = [] + markers_and_content = [] + + typedefs = [] + for type_name in sorted(constants.ANONYMOUS_STRUCTS): + s = constants.ANONYMOUS_STRUCTS[type_name] + typedefs.append( + f"using {type_name} = decltype(::{s['parent']}::{s['field_name']});" ) + markers_and_content.append(( + "// {{ ANONYMOUS_STRUCT_TYPEDEFS }}", + typedefs, + )) - autogenned_struct_definitions = [] - markers_and_content = [] + # Sort by struct name by dependency to ensure deterministic output order + sorted_struct_names = sort_structs_by_dependency(structs_to_bind_data) - typedefs = [] - for type_name in sorted(constants.ANONYMOUS_STRUCTS): - s = constants.ANONYMOUS_STRUCTS[type_name] - typedefs.append( - f"using {type_name} = decltype(::{s['parent']}::{s['field_name']});" - ) - markers_and_content.append(( - "// {{ ANONYMOUS_STRUCT_TYPEDEFS }}", - typedefs, - )) + for struct_name in sorted_struct_names: + struct_data = structs_to_bind_data[struct_name] + if struct_data.wrapped_header: + autogenned_struct_definitions.append(struct_data.wrapped_header + "\n") + else: + markers_and_content.append(( + f"// INSERT-GENERATED-{struct_data.wrap_name}-DEFINITIONS", + [ + l.definition if l.definition else "" + for l in struct_data.wrapped_fields + ], + )) + markers_and_content.append(( + "// {{ AUTOGENNED_STRUCTS_HEADER }}", + autogenned_struct_definitions, + )) - # Sort by struct name by dependency to ensure deterministic output order - sorted_struct_names = sort_structs_by_dependency(self.structs_to_bind_data) + autogenned_struct_source = [] + autogenned_struct_bindings = [] + for struct_name in sorted_struct_names: + struct_data = structs_to_bind_data[struct_name] + if struct_data.wrapped_source: + autogenned_struct_source.append(struct_data.wrapped_source + "\n") + autogenned_struct_bindings.append(struct_data.bindings) - 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") - else: - markers_and_content.append(( - f"// INSERT-GENERATED-{struct_data.wrap_name}-DEFINITIONS", - [ - l.definition if l.definition else "" - for l in struct_data.wrapped_fields - ], - )) - markers_and_content.append(( - "// {{ AUTOGENNED_STRUCTS_HEADER }}", - autogenned_struct_definitions, - )) - - autogenned_struct_source = [] - autogenned_struct_bindings = [] - 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_bindings.append(struct_data.bindings) - - markers_and_content.append(( - "// {{ AUTOGENNED_STRUCTS_SOURCE }}", - autogenned_struct_source, - )) - markers_and_content.append(( - "// {{ AUTOGENNED_STRUCTS_BINDINGS }}", - autogenned_struct_bindings, - )) - return markers_and_content + markers_and_content.append(( + "// {{ AUTOGENNED_STRUCTS_SOURCE }}", + autogenned_struct_source, + )) + markers_and_content.append(( + "// {{ AUTOGENNED_STRUCTS_BINDINGS }}", + autogenned_struct_bindings, + )) + return markers_and_content diff --git a/wasm/codegen/tests/generators_test.py b/wasm/codegen/tests/generators_test.py index 214c5da0..d9d9d289 100644 --- a/wasm/codegen/tests/generators_test.py +++ b/wasm/codegen/tests/generators_test.py @@ -19,7 +19,6 @@ from introspect import ast_nodes 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 @@ -801,7 +800,20 @@ class EnumsGeneratorTest(absltest.TestCase): def test_generate_enum_bindings(self): - generator = enums.Generator({ + expected_code = """ +EMSCRIPTEN_BINDINGS(mujoco_enums) { + 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); + +}""".strip() + + markers_and_content = enums.generate({ "TestEnum": ast_nodes.EnumDecl( name="TestEnum", declname="enum TestEnum_", @@ -818,21 +830,6 @@ class EnumsGeneratorTest(absltest.TestCase): values={}, ), }) - - expected_code = """ -EMSCRIPTEN_BINDINGS(mujoco_enums) { - 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); - -}""".strip() - - markers_and_content = generator.generate() actual_code = "\n\n".join(markers_and_content[0][1]) self.assertEqual(actual_code, expected_code)