Refactor MuJoCo WASM struct codegen.

PiperOrigin-RevId: 837530063
Change-Id: I57518b2fbb536f356a5e15839d8e17f455bccde9
This commit is contained in:
Matias Manevi
2025-11-27 07:41:42 -08:00
committed by Copybara-Service
parent 5163dfa823
commit 9ca1598b23
4 changed files with 257 additions and 467 deletions
+43 -35
View File
@@ -37,7 +37,7 @@ PRIMITIVE_TYPES: Set[str] = {
# go/keep-sorted end
}
_SKIPPED_PLUGIN_FUNCTIONS: List[str] = [
_SKIPPED_PLUGIN_FUNCTIONS = (
# go/keep-sorted start
"mj_getPluginConfig",
"mj_loadAllPluginLibraries",
@@ -59,10 +59,10 @@ _SKIPPED_PLUGIN_FUNCTIONS: List[str] = [
"mjp_registerResourceProvider",
"mjp_resourceProviderCount",
# go/keep-sorted end
]
)
# Functions that are bound as class methods
_SKIPPED_CLASS_METHODS: List[str] = [
_SKIPPED_CLASS_METHODS = (
# go/keep-sorted start
"mj_compile",
"mj_copyData",
@@ -81,15 +81,15 @@ _SKIPPED_CLASS_METHODS: List[str] = [
"mj_saveXML", # TODO(manevi): Bind this function.
"mj_saveXMLString", # TODO(manevi): Bind this function.
# go/keep-sorted end
]
)
# Omitted because not very useful
_SKIPPED_WRITABLE_ERROR: List[str] = [
_SKIPPED_WRITABLE_ERROR = (
"mj_printSchema",
]
)
# Omitted thread management functions
_SKIPPED_THREAD_FUNCTIONS: List[str] = [
_SKIPPED_THREAD_FUNCTIONS = (
# go/keep-sorted start
"mju_bindThreadPool",
"mju_defaultTask",
@@ -98,10 +98,10 @@ _SKIPPED_THREAD_FUNCTIONS: List[str] = [
"mju_threadPoolDestroy",
"mju_threadPoolEnqueue",
# go/keep-sorted end
]
)
# Omitted asset cache functions
_SKIPPED_ASSET_CACHE_FUNCTIONS: List[str] = [
_SKIPPED_ASSET_CACHE_FUNCTIONS = (
# go/keep-sorted start
"mj_clearCache",
"mj_getCache",
@@ -109,10 +109,10 @@ _SKIPPED_ASSET_CACHE_FUNCTIONS: List[str] = [
"mj_getCacheSize",
"mj_setCacheCapacity",
# go/keep-sorted end
]
)
# Omitted Virtual Filesystem (VFS) functions
_SKIPPED_VFS_FUNCTIONS: List[str] = [
_SKIPPED_VFS_FUNCTIONS = (
# go/keep-sorted start
"mj_addBufferVFS",
"mj_addFileVFS",
@@ -120,10 +120,10 @@ _SKIPPED_VFS_FUNCTIONS: List[str] = [
"mj_deleteFileVFS",
"mj_deleteVFS",
# go/keep-sorted end
]
)
# Omitted irrelevant visual functions
_SKIPPED_VISUAL_FUNCTIONS: List[str] = [
_SKIPPED_VISUAL_FUNCTIONS = (
# go/keep-sorted start
"mjv_averageCamera",
"mjv_copyData",
@@ -132,9 +132,9 @@ _SKIPPED_VISUAL_FUNCTIONS: List[str] = [
"mjv_freeScene",
"mjv_makeScene",
# go/keep-sorted end
]
)
_SKIPPED_MEMORY_FUNCTIONS: List[str] = [
_SKIPPED_MEMORY_FUNCTIONS = (
# go/keep-sorted start
"mj_freeLastXML",
"mj_freeStack",
@@ -158,9 +158,9 @@ _SKIPPED_MEMORY_FUNCTIONS: List[str] = [
"mju_warning_i",
"mju_warning_s",
# go/keep-sorted end
]
)
_SKIPPED_GETTERS_AND_SETTERS: List[str] = [
_SKIPPED_GETTERS_AND_SETTERS = (
# go/keep-sorted start
"mjs_appendFloatVec",
"mjs_appendIntVec",
@@ -179,26 +179,26 @@ _SKIPPED_GETTERS_AND_SETTERS: List[str] = [
"mjs_setStringVec",
"mjs_setUserValue",
# go/keep-sorted end
]
)
_SKIPPED_UTILITY_FUNCTIONS: List[str] = [
_SKIPPED_UTILITY_FUNCTIONS = (
# go/keep-sorted start
"mju_getXMLDependencies",
# go/keep-sorted end
]
)
# Functions that require special wrappers.
# These functions are not bound automatically but are written by hand instead.
MANUAL_WRAPPER_FUNCTIONS: List[str] = [
MANUAL_WRAPPER_FUNCTIONS = (
# go/keep-sorted start
"mj_saveLastXML",
"mj_setLengthRange",
"mju_error",
# go/keep-sorted end
]
)
# List of functions that should be skipped during the code generation process.
SKIPPED_FUNCTIONS: List[str] = (
SKIPPED_FUNCTIONS = (
_SKIPPED_CLASS_METHODS
+ _SKIPPED_THREAD_FUNCTIONS
+ _SKIPPED_MEMORY_FUNCTIONS
@@ -212,7 +212,7 @@ SKIPPED_FUNCTIONS: List[str] = (
)
# List of structs that should be skipped during the code generation process.
SKIPPED_STRUCTS: List[str] = [
SKIPPED_STRUCTS = (
# go/keep-sorted start
"mjCache",
"mjSDF",
@@ -229,14 +229,22 @@ SKIPPED_STRUCTS: List[str] = [
"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] = [
"MjvScene",
]
# Structs for which header generation is done manually.
# mjvScene is included here because buffer sizes need to be calculated based on
# introspect doc strings, which was considered a brittle unreliable solution in
# the past.
MANUAL_STRUCTS_HEADERS = (
"mjvScene",
)
# Structs for which source code generation is done manually.
MANUAL_STRUCTS_SOURCES = (
"mjData",
"mjModel",
"mjvScene",
"mjSpec",
)
# Dictionary that maps anonymous structs to their parent struct and field name.
# Anonymous structs are not defined as independent structs in the MuJoCo
@@ -262,7 +270,7 @@ STRUCTS_TO_BIND: List[str] = list(
)
# List of structs that do not have a default constructor.
NO_DEFAULT_CONSTRUCTORS: List[str] = [
NO_DEFAULT_CONSTRUCTORS = (
# go/keep-sorted start
"mjContact",
"mjSolverStat",
@@ -277,12 +285,12 @@ NO_DEFAULT_CONSTRUCTORS: List[str] = [
"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] = [
MJDATA_SIZES = (
# go/keep-sorted start
"contact",
"efc_AR",
@@ -352,7 +360,7 @@ MJDATA_SIZES: List[str] = [
"map_efc2iefc",
"map_iefc2efc",
# go/keep-sorted end
]
)
# Fields that should be entirely omitted from the bindings.
SKIPPED_FIELDS: Dict[str, List[str]] = {}
+196 -229
View File
@@ -17,7 +17,7 @@
import collections
import dataclasses
import math
from typing import Dict, List, Tuple, Union, cast
from typing import Tuple, Union, cast
from introspect import ast_nodes
from introspect import structs as introspect_structs
@@ -34,8 +34,8 @@ class WrappedFieldData:
# Line for struct field binding
binding: str = ""
# Line for struct field definition
definition: str = ""
# Line for struct field declaration
declaration: str = ""
# Used for constructor fields that are class types and need to be initialized
# with a pointer to the corresponding member in the native struct
@@ -53,27 +53,7 @@ class WrappedFieldData:
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(
def _get_property_binding(
field: ast_nodes.StructFieldDecl,
struct_wrapper_name: str,
setter: bool = False,
@@ -99,15 +79,15 @@ def _generate_field_data(
# 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")),
typename=_get_field_struct_type(f, s),
declaration=f"// {f.name} field is handled manually in template file struct declaration", # pylint: disable=line-too-long
binding=_get_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.",
declaration=f"// {f.name} field is skipped.",
binding=f"// {f.name} field is skipped.",
)
@@ -122,45 +102,32 @@ def _generate_field_data(
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),
declaration=builder.to_string(),
typename=_get_field_struct_type(f, s),
binding=_get_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.capitalize(f.type.name)} {f.name};",
typename=_get_field_struct_type(f.type),
binding=_simple_property_binding(f, w, setter=False, reference=True),
declaration=f"{common.capitalize(f.type.name)} {f.name};",
typename=_get_field_struct_type(f, s),
binding=_get_property_binding(f, w, setter=False, reference=True),
ptr_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 not anonymous_struct_name:
raise RuntimeError(
f"Anonymous struct for field {f.name} in {w} not found in"
" ANONYMOUS_STRUCTS."
)
if anonymous_struct_name in constants.STRUCTS_TO_BIND:
return WrappedFieldData(
binding=_simple_property_binding(f, w, setter=False, reference=True),
typename=anonymous_struct_name,
definition=f"{common.capitalize(anonymous_struct_name)} {f.name};",
ptr_initialization=f"{f.name}(&ptr_->{f.name})",
ptr_copy_reset=f"{f.name}.set(&ptr_->{f.name});",
is_primitive_or_fixed_size=True,
)
anonymous_struct_name = _get_field_struct_type(f, s)
return WrappedFieldData(
binding=_get_property_binding(f, w, setter=False, reference=True),
typename=anonymous_struct_name,
declaration=f"{common.capitalize(anonymous_struct_name)} {f.name};",
ptr_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):
@@ -191,19 +158,19 @@ def _generate_field_data(
)
return WrappedFieldData(
definition=builder.to_string(),
typename=_get_field_struct_type(f.type),
binding=_simple_property_binding(f, w),
declaration=builder.to_string(),
typename=_get_field_struct_type(f, s),
binding=_get_property_binding(f, w),
is_primitive_or_fixed_size=True,
)
elif inner_type_name.startswith("mj"):
return WrappedFieldData(
definition=(
declaration=(
f"std::vector<{common.capitalize(inner_type_name)}> {f.name};"
),
ptr_initialization=f"{f.name}(&ptr_->{f.name})",
typename=_get_field_struct_type(f.type),
binding=_simple_property_binding(f, w, reference=True),
typename=_get_field_struct_type(f, s),
binding=_get_property_binding(f, w, reference=True),
)
elif isinstance(f.type, ast_nodes.PointerType):
@@ -212,21 +179,10 @@ def _generate_field_data(
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":
is_dynamically_sized = bool(f.array_extent)
# Case 1: mjString
if inner_type_name == "mjString":
builder = code_builder.CodeBuilder()
with builder.function(f"mjString {f.name}() const"):
builder.line(
@@ -235,12 +191,13 @@ def _generate_field_data(
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),
declaration=builder.to_string(),
typename=_get_field_struct_type(f, s),
binding=_get_property_binding(f, w, setter=True, reference=True),
)
# Case 2: mj*Vec
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
@@ -249,47 +206,62 @@ def _generate_field_data(
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),
declaration=builder.to_string(),
typename=_get_field_struct_type(f, s),
binding=_get_property_binding(f, w, setter=False, reference=True),
)
if (
inner_type_name.startswith("mj")
# Case 3: Non-dynamically sized pointer fields to other structs.
elif (
not is_dynamically_sized
and inner_type_name not in constants.PRIMITIVE_TYPES
and not f.array_extent
and w not in constants.MANUAL_FIELDS.keys()
):
# These are wrapped as direct members of the wrapper class, initialized
# with a pointer to the corresponding member in the native struct.
ptr_field = cast(ast_nodes.PointerType, f.type)
wrapper_field_name = common.capitalize(
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),
declaration=f"{wrapper_field_name} {f.name};",
typename=_get_field_struct_type(f, s),
binding=_get_property_binding(f, w, setter=False, reference=True),
ptr_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}));"
)
# Case 4: Dynamically sized pointer fields and other pointer types that are
# exposed as emscripten::typed_memory_view.
else:
ptr_field_expr = f"ptr_->{f.name}"
array_size_str = ""
if is_dynamically_sized:
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})"
# Byte fields lack `array_extent`, so their size is defined in
# `constants.BYTE_FIELDS` based on the MuJoCo documentation.
extent = (constants.BYTE_FIELDS[f.name]["size"],)
array_size_str = parse_array_extent(extent, w, f.name)
return WrappedFieldData(
definition=builder.to_string(),
typename=_get_field_struct_type(f.type),
binding=_simple_property_binding(f, w),
)
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(
declaration=builder.to_string(),
typename=_get_field_struct_type(f, s),
binding=_get_property_binding(f, w),
)
raise RuntimeError(f"Field {f.name} from struct {w} not properly handled")
@@ -322,8 +294,8 @@ def _delete_ptr_statement(struct_name: str) -> str:
def _find_member_inits(
wrapped_fields: List[WrappedFieldData],
) -> List[WrappedFieldData]:
wrapped_fields: list[WrappedFieldData],
) -> list[WrappedFieldData]:
"""Finds the fields with ptr_initialization in the wrapped fields list."""
member_inits = []
for field in wrapped_fields:
@@ -333,7 +305,7 @@ def _find_member_inits(
def use_shallow_copy(
wrapped_fields: List[WrappedFieldData],
wrapped_fields: list[WrappedFieldData],
) -> bool:
"""Returns true if the struct fields can be shallow copied."""
for field in wrapped_fields:
@@ -344,15 +316,12 @@ def use_shallow_copy(
def build_struct_header(
struct_name: str,
wrapped_fields: List[WrappedFieldData],
wrapped_fields: list[WrappedFieldData],
):
"""Builds the C++ header file code for a struct."""
s = struct_name
w = common.capitalize(s)
if w in constants.MANUAL_STRUCTS:
return ""
if (
s not in constants.ANONYMOUS_STRUCTS
and s not in introspect_structs.STRUCTS
@@ -400,8 +369,8 @@ def build_struct_header(
# field declarations
for field in wrapped_fields:
if field.definition and field not in member_inits:
for line in field.definition.splitlines():
if field.declaration and field not in member_inits:
for line in field.declaration.splitlines():
builder.line(line)
# define private struct members
@@ -414,8 +383,8 @@ def build_struct_header(
if member_inits:
builder.public()
for f in member_inits:
if f.definition:
builder.line(f"{f.definition}")
if f.declaration:
builder.line(f"{f.declaration}")
if w == "MjData":
builder.line("mjModel *model;")
@@ -425,19 +394,9 @@ def build_struct_header(
def build_struct_source(
struct_name: str,
wrapped_fields: List[WrappedFieldData],
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.capitalize(s)
is_mjs = w.startswith("Mjs")
@@ -505,7 +464,7 @@ def build_struct_source(
def _build_struct_bindings(
struct_name: str,
wrapped_fields: List[WrappedFieldData],
wrapped_fields: list[WrappedFieldData],
):
"""Builds the C++ bindings for a struct."""
w = common.capitalize(struct_name)
@@ -536,12 +495,15 @@ def _build_struct_bindings(
if shallow_copy and not is_mjs:
builder.line(f'.function("copy", &{w}::copy, take_ownership())')
wrapped_fields.sort(key=lambda field: field.binding)
for field in wrapped_fields[:-1]:
if field.binding:
builder.line(field.binding)
if wrapped_fields:
builder.line(f"{wrapped_fields[-1].binding};")
bindings = sorted(
[f.binding for f in wrapped_fields if f.binding.startswith(".")]
)
for binding in bindings[:-1]:
builder.line(binding)
if bindings:
builder.line(f"{bindings[-1]};")
elif builder._lines:
builder._lines[-1] += ";"
return builder.to_string()
@@ -585,45 +547,6 @@ def resolve_extent(
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.capitalize(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:
@@ -644,30 +567,64 @@ def _get_anonymous_struct_field(
return target_field
def _get_field_struct_type(field_type):
def _get_field_struct_type(
field: ast_nodes.StructFieldDecl, struct_name: str
) -> str | None:
"""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
s = struct_name
w = common.capitalize(s)
if isinstance(field.type, ast_nodes.AnonymousStructDecl):
anonymous_struct_name = ""
for name, value in constants.ANONYMOUS_STRUCTS.items():
if value["parent"] == s and value["field_name"] == field.name:
anonymous_struct_name = name
break
if not anonymous_struct_name:
raise RuntimeError(
f"Anonymous struct for field {field.name} in {w} not found in"
" ANONYMOUS_STRUCTS."
)
return anonymous_struct_name
elif 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 get_introspect_struct_fields(struct_name: str):
"""Retrieves the fields of a struct from the introspect data."""
s = struct_name
if s in introspect_structs.STRUCTS:
return 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}")
return anonymous_struct.type.fields
else:
raise RuntimeError(f"Struct not found: {s}")
def sort_structs_by_dependency(
struct_wrappers: dict[str, WrappedStructData],
) -> List[str]:
struct_with_fields: dict[str, list[WrappedFieldData]],
) -> list[str]:
"""Sorts structs based on their field dependencies using topological sort.
Structs with no dependencies on other structs in the list come first.
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.
struct_with_fields: A dictionary mapping struct names to their
WrappedFieldData.
Returns:
A new list of struct names sorted by dependency.
@@ -677,25 +634,26 @@ def sort_structs_by_dependency(
"""
adj = collections.defaultdict(list)
in_degree = collections.defaultdict(int)
struct_names = struct_wrappers.keys()
struct_names = struct_with_fields.keys()
struct_set = set(struct_names)
sorted_struct_names = sorted(struct_names)
for struct_name in sorted_struct_names:
if struct_name == "mjData":
for s in sorted_struct_names:
if s == "mjData":
adj["mjModel"].append("mjData")
in_degree["mjData"] += 1
for field in struct_wrappers[struct_name].wrapped_fields:
fields = struct_with_fields[s]
for field in fields:
field_type_name = field.typename
if (
field_type_name
and field_type_name != struct_name
and field_type_name != s
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
if s not in adj[field_type_name]:
adj[field_type_name].append(s)
in_degree[s] += 1
queue = collections.deque(
[name for name in sorted_struct_names if in_degree[name] == 0]
@@ -720,14 +678,8 @@ def sort_structs_by_dependency(
)
def generate(struct_to_bind: List[str]) -> list[tuple[str, list[str]]]:
def generate(struct_to_bind: list[str]) -> 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)
markers_and_content = []
typedefs = []
for type_name in sorted(constants.ANONYMOUS_STRUCTS):
s = constants.ANONYMOUS_STRUCTS[type_name]
@@ -735,48 +687,63 @@ def generate(struct_to_bind: List[str]) -> list[tuple[str, list[str]]]:
f"using {type_name} = decltype(::{s['parent']}::{s['field_name']});"
)
# Sort by struct name by dependency to ensure deterministic output order
wrapped_structs_with_fields: dict[str, list[WrappedFieldData]] = {}
for s in struct_to_bind:
fields: list[WrappedFieldData] = []
introspect_fields = get_introspect_struct_fields(s)
for field in introspect_fields:
fields.append(_generate_field_data(field, common.capitalize(s)))
wrapped_structs_with_fields[s] = fields
dependency_sorted_struct_names = sort_structs_by_dependency(
structs_to_bind_data
wrapped_structs_with_fields
)
structs_header = []
headers = []
for s in dependency_sorted_struct_names:
struct_data = structs_to_bind_data[s]
if struct_data.wrapped_header:
structs_header.append(struct_data.wrapped_header + "\n")
else:
definitions: list[str] = []
for f in sorted(struct_data.wrapped_fields, key=lambda f: f.definition):
if f.definition:
definitions.append(f.definition)
markers_and_content.append((
f"// INSERT-GENERATED-{struct_data.wrap_name}-DEFINITIONS",
definitions,
))
fields = wrapped_structs_with_fields[s]
structs_source = []
if s not in constants.MANUAL_STRUCTS_HEADERS:
header = build_struct_header(s, fields)
headers.append(header + "\n")
sources = []
for s in dependency_sorted_struct_names:
struct_data = structs_to_bind_data[s]
if struct_data.wrapped_source:
structs_source.append(struct_data.wrapped_source + "\n")
fields = wrapped_structs_with_fields[s]
if s not in constants.MANUAL_STRUCTS_SOURCES:
source = build_struct_source(s, fields)
sources.append(source + "\n")
bindings = []
alphabetically_sorted_struct_names = sorted(
wrapped_structs_with_fields.keys()
)
for s in alphabetically_sorted_struct_names:
fields = wrapped_structs_with_fields[s]
bindings.append(_build_struct_bindings(s, fields))
struct_bindings = []
alphabetically_sorted_struct_names = sorted(structs_to_bind_data.keys())
for s in alphabetically_sorted_struct_names:
struct_data = structs_to_bind_data[s]
struct_bindings.append(struct_data.bindings)
for s in sorted(struct_to_bind):
w = common.capitalize(s)
if w.startswith("Mjs") or w == "MjSpec":
struct_bindings.append(f"emscripten::register_optional<{w}>();")
bindings.append(f"emscripten::register_optional<{w}>();")
# Combine all the markers and content into a single list
markers_and_content += [
manual_struct_field_declarations = []
for s in dependency_sorted_struct_names:
w = common.capitalize(s)
fields = wrapped_structs_with_fields[s]
if s in constants.MANUAL_STRUCTS_HEADERS:
decls: list[str] = []
for f in sorted(fields, key=lambda f: f.declaration):
if f.declaration:
decls.append(f.declaration)
manual_struct_field_declarations.append(
(f"// INSERT-GENERATED-{w}-DECLARATION", decls)
)
return [
*manual_struct_field_declarations,
("// {{ ANONYMOUS_STRUCT_TYPEDEFS }}", typedefs),
("// {{ STRUCTS_HEADER }}", structs_header),
("// {{ STRUCTS_SOURCE }}", structs_source),
("// {{ STRUCTS_BINDINGS }}", struct_bindings),
("// {{ STRUCTS_HEADER }}", headers),
("// {{ STRUCTS_SOURCE }}", sources),
("// {{ STRUCTS_BINDINGS }}", bindings),
]
return markers_and_content
+1 -1
View File
@@ -199,7 +199,7 @@ struct MjvScene {
return emscripten::val(emscripten::typed_memory_view(
6 * MjvScene::GetSumFlexFaces(), ptr_->flextexcoord));
}
// INSERT-GENERATED-MjvScene-DEFINITIONS
// INSERT-GENERATED-MjvScene-DECLARATION
private:
mjvScene* ptr_;
+17 -202
View File
@@ -203,58 +203,6 @@ class FunctionUtilsTest(absltest.TestCase):
class StructConstructorCodeBuilderTest(absltest.TestCase):
def test_constructor_code_with_default_function(self):
wrapped_structs = structs.generate_wasm_bindings(["mjLROpt"])
self.assertEqual(
wrapped_structs["mjLROpt"].wrapped_source,
"""
MjLROpt::MjLROpt(mjLROpt *ptr) : ptr_(ptr) {}
MjLROpt::~MjLROpt() {
if (owned_ && ptr_) {
delete ptr_;
}
}
MjLROpt::MjLROpt() : ptr_(new mjLROpt) {
owned_ = true;
mj_defaultLROpt(ptr_);
}
MjLROpt::MjLROpt(const MjLROpt &other) : MjLROpt() {
*ptr_ = *other.get();
}
MjLROpt& MjLROpt::operator=(const MjLROpt &other) {
if (this == &other) {
return *this;
}
*ptr_ = *other.get();
return *this;
}
std::unique_ptr<MjLROpt> MjLROpt::copy() {
return std::make_unique<MjLROpt>(*this);
}
mjLROpt* MjLROpt::get() const {
return ptr_;
}
void MjLROpt::set(mjLROpt* ptr) {
ptr_ = ptr;
}
""".strip(),
)
def test_constructor_code_without_default_function(self):
wrapped_structs = structs.generate_wasm_bindings(["mjsElement"])
self.assertEqual(
wrapped_structs["mjsElement"].wrapped_source,
"""
MjsElement::MjsElement(mjsElement *ptr) : ptr_(ptr) {}
mjsElement* MjsElement::get() const {
return ptr_;
}
void MjsElement::set(mjsElement* ptr) {
ptr_ = ptr;
}
""".strip(),
)
def test_constructor_code_with_fields_with_init(self):
field_with_init = ast_nodes.StructFieldDecl(
name="element",
@@ -354,7 +302,7 @@ class StructFieldCodeBuilderTest(absltest.TestCase):
doc="number of geoms",
)
self.assertEqual(
structs.build_primitive_type_definition(field),
structs._generate_field_data(field, "ngeom").declaration,
"""
int ngeom() const {
return ptr_->ngeom;
@@ -375,9 +323,7 @@ void set_ngeom(int value) {
array_extent=("ngeom", 4),
)
self.assertEqual(
structs.build_memory_view_definition(
field, "ptr_->ngeom * 4", "ptr_->geom_rgba"
),
structs._generate_field_data(field, "geom_rgba").declaration,
"""
emscripten::val geom_rgba() const {
return emscripten::val(emscripten::typed_memory_view(ptr_->ngeom * 4, ptr_->geom_rgba));
@@ -394,7 +340,7 @@ emscripten::val geom_rgba() const {
doc="rgba when material is omitted",
)
self.assertEqual(
structs.build_string_field_definition(field),
structs._generate_field_data(field, "MjString").declaration,
"""
mjString string_field() const {
return (ptr_ && ptr_->string_field) ? *(ptr_->string_field) : "";
@@ -416,7 +362,7 @@ void set_string_field(const mjString& value) {
doc="",
)
self.assertEqual(
structs.build_mjvec_pointer_definition(field, "mjDoubleVec"),
structs._generate_field_data(field, "mjDoubleVec").declaration,
"""
mjDoubleVec &vector_field() const {
return *(ptr_->vector_field);
@@ -432,174 +378,43 @@ mjDoubleVec &vector_field() const {
doc="",
)
self.assertEqual(
structs.build_mjvec_pointer_definition(field, "mjByteVec"),
structs._generate_field_data(field, "mjByteVec").declaration,
"""
std::vector<uint8_t> &vector_field() const {
return *(reinterpret_cast<std::vector<uint8_t>*>(ptr_->vector_field));
}""".strip(),
)
def test_simple_property_binding(self):
def test_get_property_binding(self):
field = ast_nodes.StructFieldDecl(
name="ngeom",
type=ast_nodes.ValueType(name="int"),
doc="number of geoms",
)
self.assertEqual(
structs.build_simple_property_binding(field, "MjModel"),
structs._get_property_binding(field, "MjModel"),
'.property("ngeom", &MjModel::ngeom)',
)
def test_simple_property_binding_with_setter(self):
def test_get_property_binding_with_setter(self):
field = ast_nodes.StructFieldDecl(
name="ngeom",
type=ast_nodes.ValueType(name="int"),
doc="",
)
self.assertEqual(
structs.build_simple_property_binding(field, "MjModel", True),
structs._get_property_binding(field, "MjModel", True),
'.property("ngeom", &MjModel::ngeom, &MjModel::set_ngeom)',
)
def test_simple_property_binding_with_return_value_policy_as_ref(self):
def test_get_property_binding_with_return_value_policy_as_ref(self):
field = ast_nodes.StructFieldDecl(
name="ngeom",
type=ast_nodes.ValueType(name="int"),
doc="",
)
self.assertEqual(
structs.build_simple_property_binding(
field,
"MjModel",
add_setter=True,
add_return_value_policy_as_ref=True,
),
'.property("ngeom", &MjModel::ngeom, &MjModel::set_ngeom, reference())',
)
class StructFieldCodeBuilderTest(absltest.TestCase):
def test_primitive_type_definition(self):
field = ast_nodes.StructFieldDecl(
name="ngeom",
type=ast_nodes.ValueType(name="int"),
doc="number of geoms",
)
self.assertEqual(
structs._generate_field_data(field, "ngeom").definition,
"""
int ngeom() const {
return ptr_->ngeom;
}
void set_ngeom(int value) {
ptr_->ngeom = value;
}
""".strip(),
)
def test_memory_view_definition(self):
field = ast_nodes.StructFieldDecl(
name="geom_rgba",
type=ast_nodes.PointerType(
inner_type=ast_nodes.ValueType(name="float"),
),
doc="rgba when material is omitted",
array_extent=("ngeom", 4),
)
self.assertEqual(
structs._generate_field_data(field, "geom_rgba").definition,
"""
emscripten::val geom_rgba() const {
return emscripten::val(emscripten::typed_memory_view(ptr_->ngeom * 4, ptr_->geom_rgba));
}
""".strip(),
)
def test_string_field_definition(self):
field = ast_nodes.StructFieldDecl(
name="string_field",
type=ast_nodes.PointerType(
inner_type=ast_nodes.ValueType(name="mjString"),
),
doc="rgba when material is omitted",
)
self.assertEqual(
structs._generate_field_data(field, "MjString").definition,
"""
mjString string_field() const {
return (ptr_ && ptr_->string_field) ? *(ptr_->string_field) : "";
}
void set_string_field(const mjString& value) {
if (ptr_ && ptr_->string_field) {
*(ptr_->string_field) = value;
}
}
""".strip(),
)
def test_mjvec_pointer_definition(self):
field = ast_nodes.StructFieldDecl(
name="vector_field",
type=ast_nodes.PointerType(
inner_type=ast_nodes.ValueType(name="mjDoubleVec"),
),
doc="",
)
self.assertEqual(
structs._generate_field_data(field, "mjDoubleVec").definition,
"""
mjDoubleVec &vector_field() const {
return *(ptr_->vector_field);
}""".strip(),
)
def test_mjbyte_vec_pointer_definition(self):
field = ast_nodes.StructFieldDecl(
name="vector_field",
type=ast_nodes.PointerType(
inner_type=ast_nodes.ValueType(name="mjByteVec"),
),
doc="",
)
self.assertEqual(
structs._generate_field_data(field, "mjByteVec").definition,
"""
std::vector<uint8_t> &vector_field() const {
return *(reinterpret_cast<std::vector<uint8_t>*>(ptr_->vector_field));
}""".strip(),
)
def test_simple_property_binding(self):
field = ast_nodes.StructFieldDecl(
name="ngeom",
type=ast_nodes.ValueType(name="int"),
doc="number of geoms",
)
self.assertEqual(
structs._simple_property_binding(field, "MjModel"),
'.property("ngeom", &MjModel::ngeom)',
)
def test_simple_property_binding_with_setter(self):
field = ast_nodes.StructFieldDecl(
name="ngeom",
type=ast_nodes.ValueType(name="int"),
doc="",
)
self.assertEqual(
structs._simple_property_binding(field, "MjModel", True),
'.property("ngeom", &MjModel::ngeom, &MjModel::set_ngeom)',
)
def test_simple_property_binding_with_return_value_policy_as_ref(self):
field = ast_nodes.StructFieldDecl(
name="ngeom",
type=ast_nodes.ValueType(name="int"),
doc="",
)
self.assertEqual(
structs._simple_property_binding(
structs._get_property_binding(
field,
"MjModel",
setter=True,
@@ -621,7 +436,7 @@ class StructFieldHandlerTest(absltest.TestCase):
wrapped_field_data = structs._generate_field_data(field_scalar, "MjModel")
self.assertEqual(
wrapped_field_data.definition,
wrapped_field_data.declaration,
"""
int ngeom() const {
return ptr_->ngeom;
@@ -649,7 +464,7 @@ void set_ngeom(int value) {
wrapped_field_data = structs._generate_field_data(field, "MjModel")
self.assertEqual(
wrapped_field_data.definition,
wrapped_field_data.declaration,
("""
emscripten::val geom_rgba() const {
return emscripten::val(emscripten::typed_memory_view(ptr_->ngeom * 4, ptr_->geom_rgba));
@@ -674,7 +489,7 @@ emscripten::val geom_rgba() const {
wrapped_field_data = structs._generate_field_data(field, "MjData")
self.assertEqual(
wrapped_field_data.definition,
wrapped_field_data.declaration,
("""
emscripten::val buffer() const {
return emscripten::val(emscripten::typed_memory_view(model->nbuffer, static_cast<uint8_t*>(ptr_->buffer)));
@@ -693,7 +508,7 @@ emscripten::val buffer() const {
doc="",
)
wrapped_field_data = structs._generate_field_data(field, "MjsTexture")
self.assertEqual(wrapped_field_data.definition, "MjsElement element;")
self.assertEqual(wrapped_field_data.declaration, "MjsElement element;")
self.assertEqual(
wrapped_field_data.binding,
@@ -717,7 +532,7 @@ emscripten::val buffer() const {
wrapped_field_data = structs._generate_field_data(field, "MjOption")
self.assertEqual(
wrapped_field_data.definition,
wrapped_field_data.declaration,
("""
emscripten::val gravity() const {
return emscripten::val(emscripten::typed_memory_view(3, ptr_->gravity));
@@ -741,7 +556,7 @@ emscripten::val gravity() const {
)
wrapped_field_data = structs._generate_field_data(field, "MjModel")
self.assertEqual(
wrapped_field_data.definition,
wrapped_field_data.declaration,
"""
emscripten::val multi_dim_array() const {
return emscripten::val(emscripten::typed_memory_view(12, reinterpret_cast<float*>(ptr_->multi_dim_array)));