Cleanup WASM bindings generators so they all have a similar interface

PiperOrigin-RevId: 827965451
Change-Id: I657fc4d3a480fbc62124c84fe9ddc0900a91183f
This commit is contained in:
Matija Kecman
2025-11-04 07:31:36 -08:00
committed by Copybara-Service
parent 776fc32eb4
commit 29c1a3a0a5
10 changed files with 78 additions and 109 deletions
+25 -53
View File
@@ -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())
+4 -9
View File
@@ -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,
+2
View File
@@ -11846,6 +11846,7 @@ std::optional<MjsPlugin> 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);
+6 -3
View File
@@ -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])]
+6 -3
View File
@@ -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]),
]
+4 -4
View File
@@ -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>("EmptyEnum");
"""
enum_<EmptyEnum>("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):
+17 -23
View File
@@ -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
+11 -11
View File
@@ -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
+1 -1
View File
@@ -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;")
+2 -2
View File
@@ -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__":