Cleanup WASM bindings generators so they all have a similar interface
PiperOrigin-RevId: 827965451 Change-Id: I657fc4d3a480fbc62124c84fe9ddc0900a91183f
This commit is contained in:
committed by
Copybara-Service
parent
776fc32eb4
commit
29c1a3a0a5
@@ -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())
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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])]
|
||||
|
||||
@@ -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]),
|
||||
]
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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;")
|
||||
|
||||
|
||||
@@ -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__":
|
||||
|
||||
Reference in New Issue
Block a user