diff --git a/wasm/codegen/generated/bindings.cc b/wasm/codegen/generated/bindings.cc index ca0cc9b6..c59a9278 100644 --- a/wasm/codegen/generated/bindings.cc +++ b/wasm/codegen/generated/bindings.cc @@ -771,6 +771,7 @@ EMSCRIPTEN_BINDINGS(mujoco_enums) { .value("mjSECT_CLOSED", mjSECT_CLOSED) .value("mjSECT_OPEN", mjSECT_OPEN) .value("mjSECT_FIXED", mjSECT_FIXED); + } using mjVisualGlobal = decltype(::mjVisual::global); diff --git a/wasm/codegen/generators/enums.py b/wasm/codegen/generators/enums.py index d98aa25e..398f070b 100644 --- a/wasm/codegen/generators/enums.py +++ b/wasm/codegen/generators/enums.py @@ -27,24 +27,20 @@ class Generator: def __init__(self, enums: Mapping[str, ast_nodes.EnumDecl]): self.enums = enums - def _generate_enum_binding(self, enum: ast_nodes.EnumDecl) -> str: - """Generates the Embind code for a single enum.""" - - code = f'{code_builder.INDENT}enum_<{enum.name}>("{enum.name}")' - - for value_name in enum.values: - code += f'\n{2*code_builder.INDENT}.value("{value_name}", {value_name})' - - code += ";" - return code - 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)) + 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 = "\n\n".join(code) - marker = "// {{ ENUM_BINDINGS }}" + content = builder.to_string() + marker = '// {{ ENUM_BINDINGS }}' return [(marker, [content])] diff --git a/wasm/codegen/templates/bindings.cc b/wasm/codegen/templates/bindings.cc index bae60fc2..99394700 100644 --- a/wasm/codegen/templates/bindings.cc +++ b/wasm/codegen/templates/bindings.cc @@ -155,9 +155,7 @@ EMSCRIPTEN_BINDINGS(mujoco_constants) { emscripten::function("get_mjRNDSTRING", &get_mjRNDSTRING); } -EMSCRIPTEN_BINDINGS(mujoco_enums) { // {{ ENUM_BINDINGS }} -} // {{ ANONYMOUS_STRUCT_TYPEDEFS }} diff --git a/wasm/codegen/tests/generators_test.py b/wasm/codegen/tests/generators_test.py index b68c9987..214c5da0 100644 --- a/wasm/codegen/tests/generators_test.py +++ b/wasm/codegen/tests/generators_test.py @@ -819,7 +819,9 @@ class EnumsGeneratorTest(absltest.TestCase): ), }) - expected_code = """ enum_("TestEnum") + expected_code = """ +EMSCRIPTEN_BINDINGS(mujoco_enums) { + enum_("TestEnum") .value("FIRST_VAL", FIRST_VAL) .value("SECOND_VAL", SECOND_VAL) .value("THIRD_VAL", THIRD_VAL); @@ -828,7 +830,7 @@ class EnumsGeneratorTest(absltest.TestCase): .value("ALPHA", ALPHA) .value("BETA", BETA); - enum_("EmptyEnum");""" +}""".strip() markers_and_content = generator.generate() actual_code = "\n\n".join(markers_and_content[0][1])