Refactor WASM binding directory structure

* All tests moved into wasm/codegen/tests folder
* Merged wasm/codegen/helpers/ into wasm/codegen/generators/

PiperOrigin-RevId: 829471173
Change-Id: I2dbc5d9351771817ec260c7ddf87c31e66a9eabe
This commit is contained in:
Matija Kecman
2025-11-07 09:36:27 -08:00
committed by Copybara-Service
parent 44220fcc51
commit 4e46db8903
17 changed files with 1203 additions and 1262 deletions
@@ -0,0 +1,71 @@
# Copyright 2025 DeepMind Technologies Limited
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from introspect import ast_nodes
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
class BindingBuilder:
"""Builds WASM bindings for MuJoCo."""
def __init__(
self,
template_path_cc: str,
):
with open(template_path_cc, "r") as f:
self.content_cc = f.readlines()
self.markers_and_content = []
def set_enums(self):
"""Generates and sets the 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."""
generator = structs.Generator()
self.markers_and_content += generator.generate()
return self
def set_functions(self):
"""Generates and sets the function wrappers and bindings."""
functions_to_bind: dict[str, ast_nodes.FunctionDecl] = {}
for name, func in introspect_functions.FUNCTIONS.items():
if not functions.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 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):
common.write_to_file(generated_path_cc, self.to_string())
+103
View File
@@ -0,0 +1,103 @@
# Copyright 2025 DeepMind Technologies Limited
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Helper class to build code string line by line with indentation."""
INDENT = " "
class CodeBuilder:
"""Helper class to build code string line by line with indentation."""
def __init__(self, indent_str: str = INDENT):
self._lines = []
self._indent_level = 0
self._indent_str = indent_str
def line(self, line_content: str) -> None:
"""Adds a line with indentation, special-casing "private:" and "public:"."""
indent = self._indent_str * self._indent_level
content = line_content.strip()
if content == "private:" or content == "public:":
self._lines.append(indent[:-1] + line_content)
elif content:
self._lines.append(indent + line_content)
else:
self._lines.append("")
def newline(self) -> None:
"""Adds a newline."""
self.line("")
def to_string(self) -> str:
"""Returns the complete code string."""
return "\n".join(self._lines)
class IndentBlock:
"""Helper class to manage indentation within a `with` statement."""
def __init__(self, builder: "CodeBuilder", header_line="", braces=True):
self._builder = builder
self._header_line = header_line
self._braces = braces
def __enter__(self):
line = self._header_line
if self._braces:
line += " {" if line else "{"
self._builder.line(line)
self._builder._indent_level += 1
self._line_count_enter = len(self._builder._lines)
return self._builder
def __exit__(self, exc_type, exc_val, exc_tb):
if self._builder._indent_level > 0:
self._builder._indent_level -= 1
if self._line_count_enter == len(self._builder._lines):
self._builder._lines[-1] += "}"
else:
if self._braces:
self._builder.line("}")
def block(self, header_line="", braces=True) -> IndentBlock:
"""Creates a block including braces and an optional header before the opening brace.
Use via a `with` statement.
Args:
header_line: Optional header line to add before the opening brace.
braces: Whether to include opening and closing braces.
Returns:
An IndentBlock instance that manages the indentation.
"""
return self.IndentBlock(self, header_line, braces)
def function(self, signature="") -> IndentBlock:
"""Creates a function."""
return self.block(signature)
def struct(self, name="") -> IndentBlock:
"""Creates a struct."""
return self.block(f"struct {name}")
def private(self) -> None:
"""Creates a private section."""
self.newline()
self.line("private:")
def public(self) -> None:
"""Creates a public section."""
self.newline()
self.line("public:")
+66
View File
@@ -0,0 +1,66 @@
# Copyright 2025 DeepMind Technologies Limited
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Utility functions for code generation."""
import os
def write_to_file(filepath: str, content: str) -> None:
"""Writes content to a file."""
output_dir = os.path.dirname(filepath)
try:
if output_dir:
os.makedirs(output_dir, exist_ok=True)
with open(filepath, "w") as f:
chars = f.write(content)
print(f"wrote {chars} characters to file '{filepath}'")
except IOError as e:
print(f"Error writing to output file: {filepath} - {e}")
def lowercase_first_letter(input_string: str) -> str:
"""Lowercases the first letter of a string."""
return input_string[:1].lower() + input_string[1:]
def uppercase_first_letter(input_string: str) -> str:
"""Uppercases the first letter of a string."""
return input_string[:1].upper() + input_string[1:]
def replace_lines_containing_marker(
lines: list[str],
marker: str,
content: list[str],
) -> list[str]:
"""Replaces lines containing a specific marker with new content."""
for i, line in enumerate(lines):
if marker in line:
indent = line[: len(line) - len(line.lstrip(" "))]
replacement_lines = []
for text in content:
if text.strip():
# Prepend indent to ensure the first replacement line matches the
# indentation of the marker and also ensure that text containing
# newlines is also indented correctly.
# TODO(matijak): This is working around an upstream problem, we should
# make it a precondition that content elements do not contain newlines
# and fix callers to ensure that.
replacement_lines.append(
indent + text.replace("\n", f"\n{indent}") + "\n"
)
return lines[:i] + replacement_lines + lines[i + 1 :]
return lines
+482
View File
@@ -0,0 +1,482 @@
# Copyright 2025 DeepMind Technologies Limited
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Constants used in the code generation process."""
from typing import Dict, List, Set
from introspect import structs as introspect_structs
PRIMITIVE_TYPES: Set[str] = {
# go/keep-sorted start
"char",
"double",
"float",
"int",
"mjtByte",
"mjtMeshBuiltin",
"mjtNum",
"mjtObj", # Adding this to the primitives because it is used as int,
"mjtSize",
"size_t",
"uint64_t",
"uintptr_t",
"unsigned char",
"unsigned int",
"void",
# go/keep-sorted end
}
_PLUGIN_FUNCTIONS: List[str] = [
# go/keep-sorted start
"mj_getPluginConfig",
"mj_loadAllPluginLibraries",
"mj_loadPluginLibrary",
"mjc_distance",
"mjc_getSDF",
"mjc_gradient",
"mjp_defaultDecoder",
"mjp_defaultPlugin",
"mjp_defaultResourceProvider",
"mjp_findDecoder",
"mjp_getPlugin",
"mjp_getPluginAtSlot",
"mjp_getResourceProvider",
"mjp_getResourceProviderAtSlot",
"mjp_pluginCount",
"mjp_registerDecoder",
"mjp_registerPlugin",
"mjp_registerResourceProvider",
"mjp_resourceProviderCount",
# go/keep-sorted end
]
# Functions that are bound as class methods
_CLASS_METHODS: List[str] = [
# go/keep-sorted start
"mj_compile",
"mj_copyData",
"mj_copyModel",
"mj_copySpec",
"mj_deleteData",
"mj_deleteModel",
"mj_deleteSpec",
"mj_loadXML",
"mj_makeData",
"mj_makeSpec",
"mj_parse", # TODO(manevi): Bind this function.
"mj_parseXML", # TODO(manevi): Bind this function.
"mj_parseXMLString",
"mj_recompile", # TODO(manevi): Bind this function.
"mj_saveXML", # TODO(manevi): Bind this function.
"mj_saveXMLString", # TODO(manevi): Bind this function.
# go/keep-sorted end
]
# Omitted because not very useful
_WRITABLE_ERROR: List[str] = [
"mj_printSchema",
]
# Omitted thread management functions
_THREAD_FUNCTIONS: List[str] = [
# go/keep-sorted start
"mju_bindThreadPool",
"mju_defaultTask",
"mju_taskJoin",
"mju_threadPoolCreate",
"mju_threadPoolDestroy",
"mju_threadPoolEnqueue",
# go/keep-sorted end
]
# Omitted asset cache functions
_ASSET_CACHE_FUNCTIONS: List[str] = [
# go/keep-sorted start
"mj_clearCache",
"mj_getCache",
"mj_getCacheCapacity",
"mj_getCacheSize",
"mj_setCacheCapacity",
# go/keep-sorted end
]
# Omitted Virtual Filesystem (VFS) functions
_VFS_FUNCTIONS: List[str] = [
# go/keep-sorted start
"mj_addBufferVFS",
"mj_addFileVFS",
"mj_defaultVFS",
"mj_deleteFileVFS",
"mj_deleteVFS",
# go/keep-sorted end
]
# Omitted irrelevant visual functions
_VISUAL_FUNCTIONS: List[str] = [
# go/keep-sorted start
"mjv_averageCamera",
"mjv_copyData",
"mjv_copyModel",
"mjv_defaultScene",
"mjv_freeScene",
"mjv_makeScene",
# go/keep-sorted end
]
_MEMORY_FUNCTIONS: List[str] = [
# go/keep-sorted start
"mj_freeLastXML",
"mj_freeStack",
"mj_loadModel",
"mj_markStack",
"mj_saveModel",
"mj_stackAllocByte",
"mj_stackAllocInt",
"mj_stackAllocNum",
"mj_warning",
"mjs_bodyToFrame",
"mju_boxQPmalloc",
"mju_clearHandlers",
"mju_error",
"mju_error_i",
"mju_error_s",
"mju_free",
"mju_malloc",
"mju_strncpy",
"mju_warning",
"mju_warning_i",
"mju_warning_s",
# go/keep-sorted end
]
_GETTERS_AND_SETTERS: List[str] = [
# go/keep-sorted start
"mjs_appendFloatVec",
"mjs_appendIntVec",
"mjs_appendString",
"mjs_getDouble",
"mjs_getPluginAttributes",
"mjs_getString",
"mjs_getUserValue",
"mjs_setBuffer",
"mjs_setDouble",
"mjs_setFloat",
"mjs_setInStringVec",
"mjs_setInt",
"mjs_setPluginAttributes",
"mjs_setString",
"mjs_setStringVec",
"mjs_setUserValue",
# go/keep-sorted end
]
_UTILITY_FUNCTIONS: List[str] = [
# go/keep-sorted start
"mju_getXMLDependencies",
# go/keep-sorted end
]
# 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
)
# Functions that require special wrappers to infer sizes and make additional
# validation checks. These functions are not bound automatically but are
# written by hand instead.
BOUNDCHECK_FUNCS: List[str] = [
# go/keep-sorted start
"mj_addM",
"mj_angmomMat",
"mj_applyFT",
"mj_constraintUpdate",
"mj_differentiatePos",
"mj_fullM",
"mj_geomDistance",
"mj_getState",
"mj_integratePos",
"mj_jac",
"mj_jacBody",
"mj_jacBodyCom",
"mj_jacDot",
"mj_jacGeom",
"mj_jacPointAxis",
"mj_jacSite",
"mj_jacSubtreeCom",
"mj_mulJacTVec",
"mj_mulJacVec",
"mj_mulM",
"mj_mulM2",
"mj_multiRay",
"mj_normalizeQuat",
"mj_rne",
"mj_saveLastXML",
"mj_setLengthRange",
"mj_setState",
"mj_solveM",
"mj_solveM2",
"mjd_inverseFD",
"mjd_subQuat",
"mjd_transitionFD",
"mju_L1",
"mju_add",
"mju_addScl",
"mju_addTo",
"mju_addToScl",
"mju_band2Dense",
"mju_bandMulMatVec",
"mju_boxQP",
"mju_cholFactor",
"mju_cholFactorBand",
"mju_cholSolve",
"mju_cholSolveBand",
"mju_cholUpdate",
"mju_copy",
"mju_d2n",
"mju_decodePyramid",
"mju_dense2Band",
"mju_dense2sparse",
"mju_dot",
"mju_encodePyramid",
"mju_eye",
"mju_f2n",
"mju_fill",
"mju_insertionSort",
"mju_insertionSortInt",
"mju_isZero",
"mju_mulMatMat",
"mju_mulMatMatT",
"mju_mulMatTMat",
"mju_mulMatTVec",
"mju_mulMatVec",
"mju_mulVecMatVec",
"mju_n2d",
"mju_n2f",
"mju_norm",
"mju_normalize",
"mju_printMatSparse",
"mju_scl",
"mju_sparse2dense",
"mju_sqrMatTD",
"mju_sub",
"mju_subFrom",
"mju_sum",
"mju_symmetrize",
"mju_transpose",
"mju_zero",
# go/keep-sorted end
]
# List of structs that should be skipped during the code generation process.
SKIPPED_STRUCTS: List[str] = [
# go/keep-sorted start
"mjCache",
"mjSDF",
"mjTask",
"mjThreadPool",
"mjUI",
"mjVFS",
"mjrContext",
"mjrRect",
"mjuiDef",
"mjuiItem",
"mjuiSection",
"mjuiState",
"mjuiThemeColor",
"mjuiThemeSpacing"
# go/keep-sorted end
]
# These structs require specific function calls for creation and/or deletion,
# or some of their fields need to be handled manually for now;
# making their wrapper constructors/destructors non-trivial.
MANUAL_STRUCTS: List[str] = [
"MjData",
"MjModel",
"MjvScene",
"MjSpec",
"MjVisual",
]
# Dictionary that maps anonymous structs to their parent struct and field name.
# Anonymous structs are not defined as independent structs in the MuJoCo
# codebase, but they are part of other structs. This dictionary is used to
# handle them as if they were independent structs.
ANONYMOUS_STRUCTS: Dict[str, Dict[str, str]] = {
# go/keep-sorted start
"mjVisualGlobal": {"parent": "mjVisual", "field_name": "global"},
"mjVisualHeadlight": {"parent": "mjVisual", "field_name": "headlight"},
"mjVisualMap": {"parent": "mjVisual", "field_name": "map"},
"mjVisualQuality": {"parent": "mjVisual", "field_name": "quality"},
"mjVisualRgba": {"parent": "mjVisual", "field_name": "rgba"},
"mjVisualScale": {"parent": "mjVisual", "field_name": "scale"},
# go/keep-sorted end
}
# This list is created by subtracting the skipped structs from the list of all
# structs and adding the anonymous structs.
STRUCTS_TO_BIND: List[str] = list(
set(introspect_structs.STRUCTS.keys())
.union(ANONYMOUS_STRUCTS.keys())
.difference(set(SKIPPED_STRUCTS))
)
# List of structs that do not have a default constructor.
NO_DEFAULT_CONSTRUCTORS: List[str] = [
# go/keep-sorted start
"mjContact",
"mjSolverStat",
"mjStatistic",
"mjTimerStat",
"mjWarningStat",
"mjsCompiler",
"mjsDefault",
"mjsElement",
"mjsExclude",
"mjsWrap",
"mjvGLCamera",
"mjvLight",
# go/keep-sorted end
]
# List of `mjData` fields where the array size should be obtained from other
# `mjData` members, instead of from `mjModel` members. This is typically the
# case for fields that are dynamically allocated during the simulation.
MJDATA_SIZES: List[str] = [
# go/keep-sorted start
"contact",
"efc_AR",
"efc_AR_colind",
"efc_AR_rowadr",
"efc_AR_rownnz",
"efc_D",
"efc_J",
"efc_JT",
"efc_JT_colind",
"efc_J_colind",
"efc_J_rowadr",
"efc_J_rownnz",
"efc_J_rowsuper",
"efc_KBIP",
"efc_R",
"efc_aref",
"efc_b",
"efc_diagApprox",
"efc_force",
"efc_frictionloss",
"efc_id",
"efc_island",
"efc_margin",
"efc_pos",
"efc_state",
"efc_type",
"efc_vel",
"iLDiagInv",
"iM_rowadr",
"iM_rownnz",
"iacc",
"iacc_smooth",
"iefc_D",
"iefc_J",
"iefc_JT",
"iefc_JT_colind",
"iefc_JT_rowadr",
"iefc_JT_rownnz",
"iefc_JT_rowsuper",
"iefc_J_colind",
"iefc_J_rowadr",
"iefc_J_rownnz",
"iefc_J_rowsuper",
"iefc_R",
"iefc_aref",
"iefc_force",
"iefc_frictionloss",
"iefc_id",
"iefc_state",
"iefc_type",
"ifrc_constraint",
"ifrc_smooth",
"island_dofadr",
"island_dofnum",
"island_efcadr",
"island_efcind",
"island_efcnum",
"island_idofadr",
"island_iefcadr",
"island_itreeadr",
"island_ne",
"island_nefc",
"island_nf",
"island_ntree",
"island_nv",
"map_efc2iefc",
"map_iefc2efc",
# go/keep-sorted end
]
# Fields that should be entirely omitted from the bindings.
SKIPPED_FIELDS: Dict[str, List[str]] = {}
# Fields handled manually in template file struct declaration.
MANUAL_FIELDS: Dict[str, List[str]] = {
# go/keep-sorted start
"MjData": ["solver", "timer", "warning", "contact"],
"MjModel": ["opt", "vis", "stat"],
"MjSpec": ["option", "visual", "stat", "element", "compiler"],
"MjvScene": [
# go/keep-sorted start
"camera",
"flexedge",
"flexedgeadr",
"flexedgenum",
"flexface",
"flexfaceadr",
"flexfacenum",
"flexfaceused",
"flexnormal",
"flextexcoord",
"flexvert",
"flexvertadr",
"flexvertnum",
"geomorder",
"geoms",
"lights",
"model",
"skinfacenum",
"skinnormal",
"skinvert",
"skinvertadr",
"skinvertnum",
# go/keep-sorted end
],
# go/keep-sorted end
}
# Dictionary that maps byte array fields to their corresponding size members.
# When generating the code for these fields, a specific cast to `uint8_t*` is
# required for embind. This dictionary is used to register those fields and
# their sizes.
BYTE_FIELDS: Dict[str, Dict[str, str]] = {
"buffer": {"size": "nbuffer"},
"arena": {"size": "narena"},
}
+2 -2
View File
@@ -14,11 +14,11 @@
"""Generates Embind bindings for MuJoCo enums."""
from typing import Mapping, Optional
from typing import Mapping
from introspect import ast_nodes
from wasm.codegen.helpers import code_builder
from wasm.codegen.generators import code_builder
class Generator:
+332 -6
View File
@@ -12,14 +12,340 @@
# See the License for the specific language governing permissions and
# limitations under the License.
"""Generates Embind bindings for MuJoCo functions."""
"""Helper functions for processing and generating bindings for MuJoCo functions."""
from typing import List, Mapping
from typing import List, Mapping, Set, Tuple, cast
from introspect import ast_nodes
from wasm.codegen.helpers import code_builder
from wasm.codegen.helpers import functions as function_utils
from wasm.codegen.generators import code_builder
from wasm.codegen.generators import common
from wasm.codegen.generators import constants
PRIMITIVE_TYPES = constants.PRIMITIVE_TYPES
uppercase_first_letter = common.uppercase_first_letter
def param_is_primitive_value(param: ast_nodes.FunctionParameterDecl) -> bool:
"""Checks if param is a primitive value type."""
if isinstance(param.type, ast_nodes.ValueType):
return param.type.name in PRIMITIVE_TYPES
return False
def param_is_pointer_to_primitive_value(
param: ast_nodes.FunctionParameterDecl,
) -> bool:
"""Checks if param is a pointer to a primitive value."""
return (
isinstance(param.type, ast_nodes.PointerType)
or isinstance(param.type, ast_nodes.ArrayType)
) and (
isinstance(param.type.inner_type, ast_nodes.ValueType)
and param.type.inner_type.name in PRIMITIVE_TYPES
)
def param_is_pointer_to_struct(param: ast_nodes.FunctionParameterDecl) -> bool:
"""Checks if param is a pointer to a struct."""
return (
isinstance(param.type, ast_nodes.PointerType)
or isinstance(param.type, ast_nodes.ArrayType)
) and (
isinstance(param.type.inner_type, ast_nodes.ValueType)
and param.type.inner_type.name not in PRIMITIVE_TYPES
)
def return_is_value_of_type(
func: ast_nodes.FunctionDecl, allowed_types: Set[str]
) -> bool:
"""Checks if func returns an allowed value type."""
return (
isinstance(func.return_type, ast_nodes.ValueType)
and func.return_type.name in allowed_types
)
def return_is_pointer_to_struct(func: ast_nodes.FunctionDecl) -> bool:
"""Checks if func returns a pointer to a struct."""
return (
isinstance(func.return_type, ast_nodes.PointerType)
and isinstance(func.return_type.inner_type, ast_nodes.ValueType)
and func.return_type.inner_type.name not in PRIMITIVE_TYPES
)
def return_is_pointer_to_primitive(func: ast_nodes.FunctionDecl) -> bool:
"""Checks if func returns a pointer to a primitive value."""
return (
isinstance(func.return_type, ast_nodes.PointerType)
and isinstance(func.return_type.inner_type, ast_nodes.ValueType)
and func.return_type.inner_type.name in PRIMITIVE_TYPES
)
def get_const_qualifier(func: ast_nodes.FunctionDecl) -> str:
"""Returns the const qualifier of func's return type."""
if (
isinstance(func.return_type, ast_nodes.PointerType)
and isinstance(func.return_type.inner_type, ast_nodes.ValueType)
and func.return_type.inner_type.is_const
):
return "const "
return ""
def should_be_wrapped(func: ast_nodes.FunctionDecl) -> bool:
"""Checks if a MuJoCo function needs a wrapper function."""
return (
return_is_pointer_to_primitive(func)
or return_is_pointer_to_struct(func)
or any(
param_is_pointer_to_primitive_value(param)
or isinstance(param.type, ast_nodes.ArrayType)
or param_is_pointer_to_struct(param)
for param in func.parameters
)
)
def generate_function_wrapper(func: ast_nodes.FunctionDecl) -> str:
"""Generates C++ code for a wrapper function."""
params_unpack_statements = get_params_unpack_statements(func.parameters)
wrapper_params_list = get_params_string(func.parameters)
not_nullable_params = get_params_notnullable(func.parameters)
wrapper_params = ", ".join(wrapper_params_list)
ret_type = get_compatible_return_type(func)
builder = code_builder.CodeBuilder()
with builder.function(f"{ret_type} {func.name}_wrapper({wrapper_params})"):
invoker_params_list = get_params_string_maybe_with_conversion(
func.parameters
)
invoker_params_str = ", ".join(invoker_params_list)
invoker_call = f"{func.name}({invoker_params_str})"
invoker_statement = get_compatible_return_call(func, invoker_call)
for p in not_nullable_params:
builder.line(f"CHECK_VAL({p});")
for unpack_statement in params_unpack_statements:
builder.line(unpack_statement)
builder.line(f"{invoker_statement};")
return builder.to_string()
def get_params_notnullable(
ast_params: Tuple[ast_nodes.FunctionParameterDecl, ...],
) -> List[str]:
"""Generates list of param names for checking if they aren't null/undefined."""
not_nullable_params = []
for p in ast_params:
if (
isinstance(p.type, (ast_nodes.PointerType, ast_nodes.ArrayType))
and isinstance(p.type.inner_type, ast_nodes.ValueType)
# We only check for char because others are checked in the unpacker
# and we don't want to check twice.
and p.type.inner_type.name == "char"
and not p.nullable
):
not_nullable_params.append(p.name)
return not_nullable_params
def get_params_unpack_statements(
ast_params: Tuple[ast_nodes.FunctionParameterDecl, ...],
) -> List[str]:
"""Generates C++ statements to unpack JS values for pointer/array parameters."""
params_unpack_statements = []
for p in ast_params:
if (
isinstance(p.type, (ast_nodes.PointerType, ast_nodes.ArrayType))
and isinstance(p.type.inner_type, ast_nodes.ValueType)
and p.type.inner_type.name in PRIMITIVE_TYPES
):
if p.type.inner_type.name == "char":
# param is Javascript string
continue
if p.type.inner_type.is_const:
# param is Javascript number[]
params_unpack_statements.append(
f"UNPACK_ARRAY({p.type.inner_type.name}, {p.name});"
)
else:
# param is TypedArray or a WasmBuffer
params_unpack_statements.append(
f"UNPACK_VALUE({p.type.inner_type.name}, {p.name});"
)
return params_unpack_statements
def get_params_string(
parameters: Tuple[ast_nodes.FunctionParameterDecl, ...]
) -> List[str]:
"""Generates a list of C++ parameter declarations as strings."""
result = []
for p in parameters:
if (
isinstance(p.type, ast_nodes.PointerType)
and isinstance(p.type.inner_type, ast_nodes.ValueType)
and p.type.inner_type.name not in PRIMITIVE_TYPES
):
# Pointer to struct parameters
const_qualifier = "const " if p.type.inner_type.is_const else ""
result.append(
f"{const_qualifier}{uppercase_first_letter(p.type.inner_type.name)}&"
f" {p.name}"
)
elif (
isinstance(p.type, ast_nodes.ValueType)
and p.type.name in PRIMITIVE_TYPES
):
# Primitive value parameters
const_qualifier = "const " if p.type.is_const else ""
result.append(f"{const_qualifier}{p.type} {p.name}")
elif (
isinstance(p.type, (ast_nodes.PointerType, ast_nodes.ArrayType))
and isinstance(p.type.inner_type, ast_nodes.ValueType)
and p.type.inner_type.name in PRIMITIVE_TYPES
):
# Pointer to primitive value parameters or arrays
if p.type.inner_type.name == "char":
if p.nullable:
result.append(f"const NullableString& {p.name}")
else:
result.append(f"const String& {p.name}")
elif (
p.type.inner_type.name
in ["int", "float", "double", "mjtNum", "mjtByte"]
and p.type.inner_type.is_const
):
result.append(f"const NumberArray& {p.name}")
else:
result.append(f"const val& {p.name}")
else:
# This case should ideally not be reached if AST is well-formed
# and types are categorized by the helper booleans correctly.
raise TypeError(
"Unable to generate param string. Unhandled parameter type:"
f" {p.type} for param '{p.name}'"
)
return result
def get_params_string_maybe_with_conversion(
ast_params: Tuple[ast_nodes.FunctionParameterDecl, ...],
) -> List[str]:
"""Generates C++ expressions for passing compatible params from JS to MuJoCo C-API functions."""
native_params = []
for p in ast_params:
if param_is_pointer_to_struct(p):
native_params.append(f"{p.name}.get()")
elif param_is_primitive_value(p):
native_params.append(p.name)
elif (
isinstance(p.type, (ast_nodes.PointerType, ast_nodes.ArrayType))
and isinstance(p.type.inner_type, ast_nodes.ValueType)
and p.type.inner_type.name in PRIMITIVE_TYPES
and p.type.inner_type.name != "char"
):
native_params.append(f"{p.name}_.data()")
elif (
isinstance(p.type, (ast_nodes.PointerType, ast_nodes.ArrayType))
and isinstance(p.type.inner_type, ast_nodes.ValueType)
and p.type.inner_type.name == "char"
):
const_qualifier = "const " if p.type.inner_type.is_const else ""
native_params.append(
f"{p.name}.as<{const_qualifier}std::string>().data()"
)
else:
raise TypeError(
f"Unhandled parameter type for conversion: {p.type} for param"
f" '{p.name}'"
)
return native_params
def get_compatible_return_call(
func: ast_nodes.FunctionDecl, invoker: str
) -> str:
"""Generates embind compatible return value conversion."""
if return_is_value_of_type(func, {"void"}):
return invoker
if isinstance(func.return_type, ast_nodes.PointerType) and isinstance(
func.return_type.inner_type, ast_nodes.ValueType
):
if func.return_type.inner_type.name == "char":
return f"return std::string({invoker})"
elif func.return_type.inner_type.name == "mjString":
return f"return *{invoker}"
if return_is_pointer_to_struct(func):
return get_converted_struct_to_class(func, invoker)
if return_is_value_of_type(func, PRIMITIVE_TYPES):
return f"return {invoker}"
raise RuntimeError(
"Failed to calculate return value conversion for function"
f" {func.name} that returns '{func.return_type}'"
)
def get_compatible_return_type(func: ast_nodes.FunctionDecl) -> str:
"""Creates embind compatible return type."""
if (
isinstance(func.return_type, ast_nodes.PointerType)
and isinstance(func.return_type.inner_type, ast_nodes.ValueType)
and func.return_type.inner_type.name in ["char", "mjString"]
):
return "std::string"
if (
isinstance(func.return_type, ast_nodes.PointerType)
and isinstance(func.return_type.inner_type, ast_nodes.ValueType)
and func.return_type.inner_type.name not in PRIMITIVE_TYPES
):
const_qualifier = get_const_qualifier(func)
return f"""{const_qualifier}std::optional<{uppercase_first_letter(func.return_type.inner_type.name)}>"""
if (
isinstance(func.return_type, ast_nodes.ValueType)
and func.return_type.name in PRIMITIVE_TYPES
):
return f"{func.return_type.name}"
return "val"
def get_converted_struct_to_class(
func: ast_nodes.FunctionDecl, invoker: str
) -> str:
"""Generates a C++ function invocation for a struct return-type function."""
const_qualifier = get_const_qualifier(func)
return_type = cast(ast_nodes.PointerType, func.return_type)
struct_name = cast(ast_nodes.ValueType, return_type.inner_type).name
class_constructor = uppercase_first_letter(struct_name)
return_str = f"{class_constructor}(result)"
return f"""{const_qualifier}{struct_name}* result = {invoker};
if (result == nullptr) {{
return std::nullopt;
}}
return {return_str}"""
def is_excluded_function_name(func_name: str) -> bool:
"""Checks if a function name should be excluded from direct binding."""
return (
func_name.startswith("mjr_")
or func_name.startswith("mjui_")
or func_name in constants.SKIPPED_FUNCTIONS
)
class Generator:
@@ -30,7 +356,7 @@ class Generator:
self.wrapper_bind_functions: List[ast_nodes.FunctionDecl] = []
for func in functions.values():
if function_utils.should_be_wrapped(func):
if should_be_wrapped(func):
self.wrapper_bind_functions.append(func)
else:
self.direct_bind_functions.append(func)
@@ -40,7 +366,7 @@ class Generator:
code = []
for func in self.wrapper_bind_functions:
wrapper_code = function_utils.generate_function_wrapper(func)
wrapper_code = generate_function_wrapper(func)
code.append(wrapper_code)
return "\n\n".join(code)
-106
View File
@@ -1,106 +0,0 @@
# Copyright 2025 DeepMind Technologies Limited
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from absl.testing import absltest
from introspect import ast_nodes
from wasm.codegen.generators import enums
from wasm.codegen.generators import functions
class EnumsGeneratorTest(absltest.TestCase):
def test_generate_enum_bindings(self):
generator = enums.Generator({
"TestEnum": ast_nodes.EnumDecl(
name="TestEnum",
declname="enum TestEnum_",
values={"FIRST_VAL": 0, "SECOND_VAL": 1, "THIRD_VAL": 2},
),
"AnotherEnum": ast_nodes.EnumDecl(
name="AnotherEnum",
declname="enum AnotherEnum_",
values={"ALPHA": 100, "BETA": 200},
),
"EmptyEnum": ast_nodes.EnumDecl(
name="EmptyEnum",
declname="enum EmptyEnum_",
values={},
),
})
expected_code = """ enum_<TestEnum>("TestEnum")
.value("FIRST_VAL", FIRST_VAL)
.value("SECOND_VAL", SECOND_VAL)
.value("THIRD_VAL", THIRD_VAL);
enum_<AnotherEnum>("AnotherEnum")
.value("ALPHA", ALPHA)
.value("BETA", BETA);
enum_<EmptyEnum>("EmptyEnum");"""
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):
super().setUp()
self.generator = functions.Generator({})
self.int_type = ast_nodes.ValueType(name="int")
def test_generate_function_binding_simple_case(self):
func_simple_void = ast_nodes.FunctionDecl(
name="do_nothing",
return_type=ast_nodes.ValueType(name="void"),
parameters=tuple(),
doc="doc",
)
self.assertEqual(
self.generator._generate_function_binding(func_simple_void),
'function("do_nothing", &do_nothing);\n',
)
def test_generate_direct_bindable_functions_simple_filter(self):
direct_bind = ast_nodes.FunctionDecl(
name="direct_bind",
return_type=self.int_type,
parameters=(
ast_nodes.FunctionParameterDecl(name="val", type=self.int_type),
),
doc="doc",
)
needs_wrap = ast_nodes.FunctionDecl(
name="needs_wrap",
return_type=ast_nodes.PointerType(inner_type=self.int_type),
parameters=tuple(),
doc="doc",
)
self.generator = functions.Generator({
"direct1": direct_bind,
"wrapped1": needs_wrap,
})
generated_code = self.generator._generate_direct_bindable_functions()
self.assertIn('function("direct_bind", &direct_bind);\n', generated_code)
self.assertNotIn('function("needs_wrap", &needs_wrap);\n', generated_code)
if __name__ == "__main__":
absltest.main()
+699 -12
View File
@@ -14,8 +14,701 @@
"""Generates Embind bindings for MuJoCo structs."""
from wasm.codegen.helpers import constants
from wasm.codegen.helpers import structs
import collections
import dataclasses
import math
from typing import Dict, List, Tuple, Union, cast
from introspect import ast_nodes
from introspect import structs as introspect_structs
from wasm.codegen.generators import code_builder
from wasm.codegen.generators import common
from wasm.codegen.generators import constants
@dataclasses.dataclass
class WrappedFieldData:
"""Data class for struct field definition and binding."""
# Line for struct field binding
binding: str = ""
# Line for struct field definition
definition: str = ""
# Initialization code for fields that require it
initialization: str = ""
# Statement to reset the inner pointer when copying the field
ptr_copy_reset: str = ""
# Whether the field is a primitive or fixed size
is_primitive_or_fixed_size: bool = False
# Underlying type of the field. If non-empty, used to determine the order in
# which structs are written in the bindings.h file.
typename: str = ""
@dataclasses.dataclass
class WrappedStructData:
"""Data class for struct wrapper definition and binding."""
# Name of wrapper struct
wrap_name: str
# List of WrappedFieldData for this struct
wrapped_fields: List[WrappedFieldData]
# Struct header code
wrapped_header: str
# Struct source code
wrapped_source: str
# Struct bindings code
bindings: str = ""
def _simple_property_binding(
field: ast_nodes.StructFieldDecl,
struct_wrapper_name: str,
setter: bool = False,
reference: bool = False,
) -> str:
"""Builds the C++ code for a simple property binding."""
f = field
w = struct_wrapper_name
setter_txt = f", &{w}::set_{f.name}" if setter else ""
reference_txt = ", reference()" if reference else ""
return f'.property("{f.name}", &{w}::{f.name}{setter_txt}{reference_txt})'
def _generate_field_data(
field: ast_nodes.StructFieldDecl, struct_wrapper_name: str
) -> WrappedFieldData:
"""Generates the C++ definition and binding code for the struct field."""
f = field
w = struct_wrapper_name
s = common.lowercase_first_letter(w)
if f.name in constants.MANUAL_FIELDS.get(w, []):
# Note: Manually handled MjModel fields are special cased so that a
# by-reference embind return value policy is used.
return WrappedFieldData(
typename=_get_field_struct_type(f.type),
definition=f"// {f.name} field is handled manually in template file struct declaration", # pylint: disable=line-too-long
binding=_simple_property_binding(f, w, reference=(w == "MjModel")),
)
if f.name in constants.SKIPPED_FIELDS.get(w, []):
return WrappedFieldData(
typename="",
definition=f"// {f.name} field is skipped.",
binding=f"// {f.name} field is skipped.",
)
if isinstance(f.type, ast_nodes.ValueType) and (
f.type.name in constants.PRIMITIVE_TYPES or f.type.name.startswith("mjt")
):
builder = code_builder.CodeBuilder()
with builder.function(f"{f.type.name} {f.name}() const"):
builder.line(f"return ptr_->{f.name};")
with builder.function(f"void set_{f.name}({f.type.name} value)"):
builder.line(f"ptr_->{f.name} = value;")
return WrappedFieldData(
definition=builder.to_string(),
typename=_get_field_struct_type(f.type),
binding=_simple_property_binding(f, w, setter=True, reference=True),
is_primitive_or_fixed_size=True,
)
elif isinstance(f.type, ast_nodes.ValueType) and f.type.name.startswith("mj"):
return WrappedFieldData(
definition=f"{common.uppercase_first_letter(f.type.name)} {f.name};",
typename=_get_field_struct_type(f.type),
binding=_simple_property_binding(f, w, setter=False, reference=True),
initialization=f", {f.name}(&ptr_->{f.name})",
ptr_copy_reset=f"{f.name}.set(&ptr_->{f.name});",
is_primitive_or_fixed_size=True,
)
elif isinstance(f.type, ast_nodes.AnonymousStructDecl):
anonymous_struct_name = ""
for name, value in constants.ANONYMOUS_STRUCTS.items():
if value["parent"] == s and value["field_name"] == f.name:
anonymous_struct_name = name
break
if anonymous_struct_name in constants.STRUCTS_TO_BIND:
return WrappedFieldData(
binding=_simple_property_binding(f, w, setter=False, reference=True),
typename=_get_field_struct_type(f.type),
initialization=f", {f.name}(&ptr_->{f.name})",
ptr_copy_reset=f"{f.name}.set(&ptr_->{f.name});",
is_primitive_or_fixed_size=True,
)
elif isinstance(f.type, ast_nodes.ArrayType):
inner_type = f.type.inner_type
size = math.prod(f.type.extents)
if (
isinstance(inner_type, ast_nodes.ValueType)
and inner_type.name in constants.PRIMITIVE_TYPES
):
ptr_expr = f"ptr_->{f.name}"
if len(f.type.extents) > 1:
# for multi-dimensional arrays, we need to cast the field
# to a pointer, so embind can correctly interpret the memory
# view
ptr_expr = f"reinterpret_cast<{inner_type.name}*>({ptr_expr})"
builder = code_builder.CodeBuilder()
with builder.function(f"emscripten::val {f.name}() const"):
builder.line(
"return"
f" emscripten::val(emscripten::typed_memory_view({str(size)},"
f" {ptr_expr}));"
)
return WrappedFieldData(
definition=builder.to_string(),
typename=_get_field_struct_type(f.type),
binding=_simple_property_binding(f, w),
is_primitive_or_fixed_size=True,
)
elif isinstance(f.type, ast_nodes.PointerType):
inner_type_name = (
f.type.inner_type.name
if isinstance(f.type.inner_type, ast_nodes.ValueType)
else ""
)
ptr_field_expr = f"ptr_->{f.name}"
array_size_str = ""
if f.array_extent:
array_size_str = parse_array_extent(f.array_extent, w, f.name)
elif f.name in constants.BYTE_FIELDS.keys():
# for byte fields, we need to cast the pointer to uint8_t*
# so embind can correctly interpret the memory view
ptr_field_expr = f"static_cast<uint8_t*>({ptr_field_expr})"
# for these byte fields, there is no array_extent, so we add the size of
# in the config file based in the documentation
extent = (constants.BYTE_FIELDS[f.name]["size"],)
array_size_str = parse_array_extent(extent, w, f.name)
elif inner_type_name == "mjString":
builder = code_builder.CodeBuilder()
with builder.function(f"mjString {f.name}() const"):
builder.line(
f'return (ptr_ && ptr_->{f.name}) ? *(ptr_->{f.name}) : "";'
)
with builder.function(f"void set_{f.name}(const mjString& value)"):
with builder.block(f"if (ptr_ && ptr_->{f.name})"):
builder.line(f"*(ptr_->{f.name}) = value;")
return WrappedFieldData(
definition=builder.to_string(),
typename=_get_field_struct_type(f.type),
binding=_simple_property_binding(f, w, setter=True, reference=True),
)
elif inner_type_name.startswith("mj") and inner_type_name.endswith("Vec"):
ptr_field_expr_vec = f"*(ptr_->{f.name})"
vector_type = inner_type_name
if vector_type == "mjByteVec":
vector_type = "std::vector<uint8_t>"
ptr_field_expr_vec = (
f"*(reinterpret_cast<std::vector<uint8_t>*>(ptr_->{f.name}))"
)
builder = code_builder.CodeBuilder()
with builder.function(f"{vector_type} &{f.name}() const"):
builder.line(f"return {ptr_field_expr_vec};")
return WrappedFieldData(
definition=builder.to_string(),
typename=_get_field_struct_type(f.type),
binding=_simple_property_binding(f, w, setter=False, reference=True),
)
if (
inner_type_name.startswith("mj")
and inner_type_name not in constants.PRIMITIVE_TYPES
and not f.array_extent
and w not in constants.MANUAL_FIELDS.keys()
):
ptr_field = cast(ast_nodes.PointerType, f.type)
wrapper_field_name = common.uppercase_first_letter(
cast(ast_nodes.ValueType, ptr_field.inner_type).name
)
return WrappedFieldData(
definition=f"{wrapper_field_name} {f.name};",
typename=_get_field_struct_type(f.type),
binding=_simple_property_binding(f, w, setter=False, reference=True),
initialization=f", {f.name}(ptr_->{f.name})",
)
builder = code_builder.CodeBuilder()
with builder.function(f"emscripten::val {f.name}() const"):
builder.line(
"return"
f" emscripten::val(emscripten::typed_memory_view({array_size_str},"
f" {ptr_field_expr}));"
)
return WrappedFieldData(
definition=builder.to_string(),
typename=_get_field_struct_type(f.type),
binding=_simple_property_binding(f, w),
)
# SHOULD NOT OCCUR
print("Error: field {f.name} not properly handled")
return WrappedFieldData(
definition=f"// Error: field {f.name} not properly handled.",
typename=_get_field_struct_type(f.type),
binding=f"// Error: field {f.name} not properly handled.",
)
def _has_nested_wrapper_members(struct_info: ast_nodes.StructDecl) -> bool:
"""Checks if the struct contains other wrapped structs as direct members."""
for field in struct_info.fields:
member_type = cast(ast_nodes.StructFieldDecl, field).type
if isinstance(member_type, (ast_nodes.ArrayType, ast_nodes.PointerType)):
member_type = member_type.inner_type
if isinstance(member_type, ast_nodes.ValueType):
return member_type.name in constants.STRUCTS_TO_BIND
return False
def _build_struct_header_internal(
struct_name: str,
wrapped_fields: List[WrappedFieldData],
fields_with_init: List[WrappedFieldData],
is_mjs: bool = False,
):
"""Builds the C++ header file code for a struct."""
s = struct_name
w = common.uppercase_first_letter(s)
shallow_copy = use_shallow_copy(wrapped_fields)
builder = code_builder.CodeBuilder()
with builder.struct(f"{w}"):
builder.line(f"explicit {w}({s} *ptr);")
builder.line(f"~{w}();")
if not is_mjs:
builder.line(f"{w}();")
if shallow_copy and not is_mjs:
builder.line(f"{w}(const {w} &);")
builder.line(f"{w} &operator=(const {w} &);")
builder.line(f"std::unique_ptr<{w}> copy();")
builder.line(f"{s}* get() const;")
builder.line(f"void set({s}* ptr);")
for field in wrapped_fields:
if field.definition and field not in fields_with_init:
for line in field.definition.splitlines():
builder.line(line)
builder.private()
builder.line(f"{s}* ptr_;")
if not is_mjs:
builder.line("bool owned_ = false;")
if is_mjs and fields_with_init:
builder.public()
for field in fields_with_init:
if field.definition:
builder.line(f"{field.definition}")
return builder.to_string() + ";"
def _default_function_statement(struct_name: str) -> str:
"""Returns the default function name for the given struct."""
if struct_name == "mjvGeom":
f = "mjv_initGeom(ptr_, mjGEOM_NONE, nullptr, nullptr, nullptr, nullptr);"
return f
elif struct_name in constants.ANONYMOUS_STRUCTS.keys():
return ""
elif struct_name in constants.NO_DEFAULT_CONSTRUCTORS:
return ""
elif struct_name.startswith("mjs"):
return f"mjs_default{struct_name.removeprefix('mjs')}(ptr_);"
elif struct_name.startswith("mjv"):
return f"mjv_default{struct_name.removeprefix('mjv')}(ptr_);"
elif struct_name.startswith("mj"):
return f"mj_default{struct_name.removeprefix('mj')}(ptr_);"
return ""
def _delete_ptr_statement(struct_name: str) -> str:
"""Returns the delete function name for the given struct."""
if struct_name == "mjVFS":
return "mj_deleteVFS(ptr_);"
return "delete ptr_;"
def _find_fields_with_init(
wrapped_fields: List[WrappedFieldData],
) -> List[WrappedFieldData]:
"""Finds the fields with initialization in the wrapped fields list."""
fields_with_init = []
for field in wrapped_fields:
if field.initialization:
fields_with_init.append(field)
return fields_with_init
def use_shallow_copy(
wrapped_fields: List[WrappedFieldData],
) -> bool:
"""Returns true if the struct fields can be shallow copied."""
for field in wrapped_fields:
if not field.is_primitive_or_fixed_size:
return False
return True
def build_struct_header(
struct_name: str,
wrapped_fields: List[WrappedFieldData],
):
"""Builds the C++ header file code for a struct."""
struct_info = introspect_structs.STRUCTS.get(struct_name)
if struct_name.startswith("mjs"):
fields_with_init = _find_fields_with_init(wrapped_fields)
return _build_struct_header_internal(
struct_name,
wrapped_fields,
fields_with_init,
is_mjs=True,
)
is_anonymous_struct = struct_name in constants.ANONYMOUS_STRUCTS.keys()
is_hardcoded_wrapper_struct = (
common.uppercase_first_letter(struct_name) in constants.MANUAL_STRUCTS
)
if (
not is_hardcoded_wrapper_struct
and struct_info
and not _has_nested_wrapper_members(struct_info)
or is_anonymous_struct
):
return _build_struct_header_internal(
struct_name, wrapped_fields, [], is_mjs=False
)
return ""
def build_struct_source(
struct_name: str,
wrapped_fields: List[WrappedFieldData],
):
"""Builds the C++ .cc file code for a struct."""
# These structs require specific function calls for creation and/or deletion
# which, for now, are hardcoded in the template file.
if struct_name in [
"mjData",
"mjModel",
"mjvScene",
"mjSpec",
]:
return ""
s = struct_name
w = common.uppercase_first_letter(s)
is_mjs = w.startswith("Mjs")
fields_with_init = _find_fields_with_init(wrapped_fields)
shallow_copy = use_shallow_copy(wrapped_fields)
fields_init = ""
if fields_with_init:
fields_init = "".join(
field_with_init.initialization for field_with_init in fields_with_init
)
builder = code_builder.CodeBuilder()
# constructor passing native ptr
with builder.function(f"{w}::{w}({s} *ptr) : ptr_(ptr){fields_init}"):
pass
# destructor
with builder.function(f"{w}::~{w}()"):
if not is_mjs:
with builder.block("if (owned_ && ptr_)"):
delete_ptr = _delete_ptr_statement(s)
builder.line(delete_ptr)
if not is_mjs:
# default constructor
with builder.function(f"{w}::{w}() : ptr_(new {s}){fields_init}"):
builder.line("owned_ = true;")
default_func = _default_function_statement(s)
if default_func:
builder.line(default_func)
if shallow_copy and not is_mjs:
# copy constructor
with builder.function(f"{w}::{w}(const {w} &other) : {w}()"):
builder.line("*ptr_ = *other.get();")
for field_with_init in fields_with_init:
if field_with_init.ptr_copy_reset is not None:
builder.line(field_with_init.ptr_copy_reset)
# assignment operator
with builder.function(f"{w}& {w}::operator=(const {w} &other)"):
with builder.block("if (this == &other)"):
builder.line("return *this;")
builder.line("*ptr_ = *other.get();")
for field_with_init in fields_with_init:
if field_with_init.ptr_copy_reset is not None:
builder.line(field_with_init.ptr_copy_reset)
builder.line("return *this;")
# explicit copy function
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;")
return builder.to_string()
def _build_struct_bindings(
struct_name: str,
wrapped_fields: List[WrappedFieldData],
):
"""Builds the C++ bindings for a struct."""
w = common.uppercase_first_letter(struct_name)
is_mjs = w.startswith("Mjs")
builder = code_builder.CodeBuilder()
with builder.block(
header_line=f'emscripten::class_<{w}>("{w}")', braces=False
):
if w == "MjData":
builder.line(".constructor<MjModel *>()")
builder.line(".constructor<const MjModel &, const MjData &>()")
elif w == "MjModel":
builder.line(
'.class_function("loadFromXML", &loadFromXML, take_ownership())'
)
builder.line(".constructor<const MjModel &>()")
elif w == "MjSpec":
builder.line(".constructor<const MjSpec &>()")
elif w == "MjvScene":
builder.line(".constructor<MjModel *, int>()")
builder.line(".constructor<>()")
elif not is_mjs:
builder.line(".constructor<>()")
shallow_copy = use_shallow_copy(wrapped_fields)
if shallow_copy and not is_mjs:
builder.line(f'.function("copy", &{w}::copy, take_ownership())')
for field in wrapped_fields[:-1]:
if field.binding:
builder.line(field.binding)
if wrapped_fields:
builder.line(f"{wrapped_fields[-1].binding};")
return builder.to_string()
def parse_array_extent(
extents: Tuple[Union[str, int], ...], wrapper_name: str, field_name: str
) -> str:
"""Parses the array extent of a field, returning a string representing the resolved extents."""
if not extents:
return ""
return " * ".join(
resolve_extent(extent, wrapper_name, field_name) for extent in extents
)
def resolve_extent(
extent: Union[str, int], wrapper_name: str, field_name: str
) -> str:
"""Resolves the extent of an array, handling integers and references to other struct fields.
Args:
extent: The extent to resolve, can be an int or a string referencing a
field.
wrapper_name: The name of the struct wrapper.
field_name: The name of the field being processed.
Returns:
A string representing the resolved extent, either as a number or a field
reference.
"""
if isinstance(extent, int):
return str(extent)
# if starts with mj, it's a mujoco constant,
# so we don't need to get a parent struct ptr
if extent.startswith("mj"):
return str(extent)
if wrapper_name == "MjData" and field_name not in constants.MJDATA_SIZES:
var_name = "model"
else:
var_name = "ptr_"
return f"{var_name}->{extent}"
def generate_wasm_bindings(
structs_to_bind: List[str],
) -> Dict[str, WrappedStructData]:
"""Generates WASM bindings for MuJoCo structs."""
wrapped_structs: Dict[str, WrappedStructData] = {}
for struct_name in structs_to_bind:
s = struct_name
w = common.uppercase_first_letter(s)
if s in introspect_structs.STRUCTS:
struct_fields = introspect_structs.STRUCTS[s].fields
elif s in constants.ANONYMOUS_STRUCTS:
anonymous_struct = _get_anonymous_struct_field(s)
if not anonymous_struct or not isinstance(
anonymous_struct.type, ast_nodes.AnonymousStructDecl
):
raise RuntimeError(f"Anonymous struct not found: {s}")
struct_fields = anonymous_struct.type.fields
else:
raise RuntimeError(f"Struct not found: {s}")
wrapped_fields: List[WrappedFieldData] = []
for field in struct_fields:
wrapped_fields.append(_generate_field_data(field, w))
wrap_data = WrappedStructData(
wrap_name=w,
wrapped_fields=wrapped_fields,
wrapped_header=build_struct_header(s, wrapped_fields),
wrapped_source=build_struct_source(s, wrapped_fields),
bindings=_build_struct_bindings(s, wrapped_fields),
)
wrapped_structs[s] = wrap_data
return wrapped_structs
def _get_anonymous_struct_field(
anonymous_structs_key: str,
) -> ast_nodes.StructFieldDecl | None:
"""Looks up the given key in the anonymous_structs dict and generates bindings for its fields."""
info = constants.ANONYMOUS_STRUCTS[anonymous_structs_key]
parent_decl = introspect_structs.STRUCTS[info["parent"]]
target_field = next(
(
f
for f in parent_decl.fields
if hasattr(f, "name")
and f.name == info["field_name"]
and hasattr(f, "type")
and isinstance(f.type, ast_nodes.AnonymousStructDecl)
),
None,
)
return target_field
def _get_field_struct_type(field_type):
"""Extracts the base struct name if the field type is a struct or pointer to a struct."""
if isinstance(field_type, ast_nodes.ValueType):
return field_type.name
if isinstance(field_type, ast_nodes.PointerType):
if isinstance(field_type.inner_type, ast_nodes.ValueType):
return field_type.inner_type.name
return None
def sort_structs_by_dependency(
struct_wrappers: dict[str, WrappedStructData],
) -> List[str]:
"""Sorts structs based on their field dependencies using topological sort.
Structs with no dependencies on other structs in the list come first.
Struct A has a dependency on struct B if struct A has a field where the
underlying_type is B. Note that this definition is stricter than the C++
struct dependency criterion where forward declarations can be used to
eliminate dependencies A and B if A only has a pointer to B.
Args:
struct_wrappers: A dictionary mapping struct names to their
WrappedStructData.
Returns:
A new list of struct names sorted by dependency.
Raises:
RuntimeError: If a cyclic dependency is detected.
"""
adj = collections.defaultdict(list)
in_degree = collections.defaultdict(int)
struct_names = struct_wrappers.keys()
struct_set = set(struct_names)
sorted_struct_names = sorted(struct_names)
for struct_name in sorted_struct_names:
for field in struct_wrappers[struct_name].wrapped_fields:
field_type_name = field.typename
if (
field_type_name
and field_type_name != struct_name
and field_type_name in struct_set
):
if struct_name not in adj[field_type_name]:
adj[field_type_name].append(struct_name)
in_degree[struct_name] += 1
queue = collections.deque(
[name for name in sorted_struct_names if in_degree[name] == 0]
)
sorted_list = []
while queue:
u = queue.popleft()
sorted_list.append(u)
for v in adj[u]:
in_degree[v] -= 1
if in_degree[v] == 0:
queue.append(v)
if len(sorted_list) == len(struct_names):
return sorted_list
else:
remaining = set(struct_names) - set(sorted_list)
raise RuntimeError(
"Cycle detected in struct dependencies, involving: "
f"{', '.join(sorted(list(remaining)))}"
)
class Generator:
@@ -26,7 +719,7 @@ class Generator:
# 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(
self.structs_to_bind_data = generate_wasm_bindings(
constants.STRUCTS_TO_BIND
)
@@ -34,16 +727,12 @@ class Generator:
markers_and_content = []
# Sort by struct name by dependency to ensure deterministic output order
sorted_struct_names = structs.sort_structs_by_dependency(
self.structs_to_bind_data
)
sorted_struct_names = sort_structs_by_dependency(self.structs_to_bind_data)
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"
)
autogenned_struct_definitions.append(struct_data.wrapped_header + "\n")
else:
markers_and_content.append((
f"// INSERT-GENERATED-{struct_data.wrap_name}-DEFINITIONS",
@@ -62,9 +751,7 @@ class Generator:
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_source.append(struct_data.wrapped_source + "\n")
autogenned_struct_bindings.append(struct_data.bindings)
markers_and_content.append((