Various minor cleanups to WASM code generation

* Add newline function to code_builder
* Remove type aliases which made jump to definition less ergonomic
* Add types to constants

PiperOrigin-RevId: 827863199
Change-Id: Ib4390d398926466ecce491a3d0ebe308ed5e0cbe
This commit is contained in:
Matija Kecman
2025-11-04 02:06:05 -08:00
committed by Copybara-Service
parent 10130297c0
commit c22d94e470
12 changed files with 190 additions and 218 deletions
+38 -45
View File
@@ -22,13 +22,6 @@ from wasm.codegen.helpers import constants
from wasm.codegen.helpers import struct_field_code_builder
from wasm.codegen.helpers import structs_wrappers_data
AnonymousStructDecl = ast_nodes.AnonymousStructDecl
ArrayType = ast_nodes.ArrayType
PointerType = ast_nodes.PointerType
StructFieldDecl = ast_nodes.StructFieldDecl
ValueType = ast_nodes.ValueType
WrappedFieldData = structs_wrappers_data.WrappedFieldData
debug_print = common.debug_print
@@ -37,7 +30,7 @@ class StructFieldHandler:
def __init__(
self,
field: StructFieldDecl,
field: ast_nodes.StructFieldDecl,
struct_wrapper_name: str,
):
self.field = field
@@ -53,27 +46,27 @@ class StructFieldHandler:
)
)
def generate(self) -> WrappedFieldData:
def generate(self) -> structs_wrappers_data.WrappedFieldData:
"""Generates the C++ definition and binding code for the struct field."""
field_type = self.field.type
if isinstance(field_type, ValueType) and (
if isinstance(field_type, ast_nodes.ValueType) and (
field_type.name in constants.PRIMITIVE_TYPES
or field_type.name.startswith("mjt")
):
return self._handle_primitive()
elif isinstance(field_type, PointerType):
elif isinstance(field_type, ast_nodes.PointerType):
return self._handle_pointer()
elif isinstance(field_type, ArrayType):
elif isinstance(field_type, ast_nodes.ArrayType):
return self._handle_array()
elif isinstance(field_type, ValueType) and field_type.name.startswith("mj"):
elif isinstance(field_type, ast_nodes.ValueType) and field_type.name.startswith("mj"):
return self._handle_mj_struct()
elif isinstance(field_type, AnonymousStructDecl):
elif isinstance(field_type, ast_nodes.AnonymousStructDecl):
return self._handle_anonymous_struct()
return self._undefined()
def _handle_primitive(self) -> WrappedFieldData:
def _handle_primitive(self) -> structs_wrappers_data.WrappedFieldData:
"""Handles the generation of C++ definition and binding code for primitive fields."""
return WrappedFieldData(
return structs_wrappers_data.WrappedFieldData(
definition=(
struct_field_code_builder.build_primitive_type_definition(
self.field
@@ -88,17 +81,17 @@ class StructFieldHandler:
is_primitive_or_fixed_size=True,
)
def _handle_pointer(self) -> WrappedFieldData:
def _handle_pointer(self) -> structs_wrappers_data.WrappedFieldData:
"""Handles the generation of C++ definition and binding code for pointer fields."""
if not isinstance(self.field.type, PointerType):
if not isinstance(self.field.type, ast_nodes.PointerType):
raise ValueError(
f"Expected PointerType, got {type(self.field.type)} for field"
f" {self.field.name}"
)
field_type: PointerType = self.field.type
field_type: ast_nodes.PointerType = self.field.type
inner_type_name = (
field_type.inner_type.name
if isinstance(field_type.inner_type, ValueType)
if isinstance(field_type.inner_type, ast_nodes.ValueType)
else ""
)
ptr_field_expr = f"ptr_->{self.field.name}"
@@ -111,9 +104,7 @@ class StructFieldHandler:
elif self.field.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})"
)
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[self.field.name]["size"],)
@@ -121,7 +112,7 @@ class StructFieldHandler:
extent, self.struct_wrapper_name, self.field.name
)
elif inner_type_name == "mjString":
return WrappedFieldData(
return structs_wrappers_data.WrappedFieldData(
definition=struct_field_code_builder.build_string_field_definition(
self.field
),
@@ -133,7 +124,7 @@ class StructFieldHandler:
),
)
elif inner_type_name.startswith("mj") and inner_type_name.endswith("Vec"):
return WrappedFieldData(
return structs_wrappers_data.WrappedFieldData(
definition=struct_field_code_builder.build_mjvec_pointer_definition(
self.field, inner_type_name
),
@@ -164,11 +155,11 @@ class StructFieldHandler:
and self.struct_wrapper_name
not in constants.MANUALLY_ADDED_FIELDS_FROM_TEMPLATE.keys()
):
ptr_field = cast(PointerType, self.field.type)
ptr_field = cast(ast_nodes.PointerType, self.field.type)
wrapper_field_name = common.uppercase_first_letter(
cast(ValueType, ptr_field.inner_type).name
cast(ast_nodes.ValueType, ptr_field.inner_type).name
)
return WrappedFieldData(
return structs_wrappers_data.WrappedFieldData(
definition=f"{wrapper_field_name} {self.field.name};",
binding=struct_field_code_builder.build_simple_property_binding(
self.field,
@@ -185,7 +176,7 @@ class StructFieldHandler:
)
return self._get_manual_definition(comment_type="complex pointer field")
return WrappedFieldData(
return structs_wrappers_data.WrappedFieldData(
definition=(
struct_field_code_builder.build_memory_view_definition(
self.field, array_size_str, ptr_field_expr
@@ -194,10 +185,10 @@ class StructFieldHandler:
binding=self.simple_property_binding,
)
def _handle_array(self) -> WrappedFieldData:
def _handle_array(self) -> structs_wrappers_data.WrappedFieldData:
"""Handles the generation of C++ definition and binding code for array fields."""
field_type = self.field.type
if not isinstance(field_type, ArrayType):
if not isinstance(field_type, ast_nodes.ArrayType):
raise ValueError(
f"Expected ArrayType, got {type(field_type)} for field"
f" {self.field.name}"
@@ -205,7 +196,7 @@ class StructFieldHandler:
inner_type = field_type.inner_type
size = math.prod(field_type.extents)
if isinstance(inner_type, ValueType):
if isinstance(inner_type, ast_nodes.ValueType):
if inner_type.name in constants.PRIMITIVE_TYPES:
ptr_expr = f"ptr_->{self.field.name}"
if len(field_type.extents) > 1:
@@ -213,7 +204,7 @@ class StructFieldHandler:
# to a pointer, so embind can correctly interpret the memory
# view
ptr_expr = f"reinterpret_cast<{inner_type.name}*>({ptr_expr})"
return WrappedFieldData(
return structs_wrappers_data.WrappedFieldData(
definition=(
struct_field_code_builder.build_memory_view_definition(
self.field, str(size), ptr_expr
@@ -229,17 +220,17 @@ class StructFieldHandler:
return self._get_manual_definition(comment_type="array field")
debug_print(f"\tNOT IMPLEMENTED ARRAY field: {self.field.name}")
return WrappedFieldData(
return structs_wrappers_data.WrappedFieldData(
definition=(
f"// TODO: NOT IMPLEMENTED ARRAY wrapper for {self.field.name}"
),
binding=f"// TODO: NOT IMPLEMENTED ARRAY binding for {self.field.name}",
)
def _handle_mj_struct(self) -> WrappedFieldData:
def _handle_mj_struct(self) -> structs_wrappers_data.WrappedFieldData:
"""Handles the generation of C++ definition and binding code for mj struct fields."""
if (
isinstance(self.field.type, ValueType)
isinstance(self.field.type, ast_nodes.ValueType)
and self.field.name not in self.manually_added_fields
and self.field.type.name in constants.STRUCTS_TO_BIND
):
@@ -249,7 +240,7 @@ class StructFieldHandler:
if self.struct_wrapper_name not in constants.HARDCODED_WRAPPER_STRUCTS:
wrapper_field_name = common.uppercase_first_letter(self.field.type.name)
definition = f"{wrapper_field_name} {self.field.name};"
return WrappedFieldData(
return structs_wrappers_data.WrappedFieldData(
definition=definition,
binding=struct_field_code_builder.build_simple_property_binding(
self.field,
@@ -263,7 +254,7 @@ class StructFieldHandler:
)
return self._get_manual_definition(comment_type="struct field")
def _handle_anonymous_struct(self) -> WrappedFieldData:
def _handle_anonymous_struct(self) -> structs_wrappers_data.WrappedFieldData:
"""Handles the generation of C++ definition and binding code for anonymous struct fields."""
anonymous_struct_name = ""
@@ -277,11 +268,11 @@ class StructFieldHandler:
break
if (
isinstance(self.field.type, AnonymousStructDecl)
isinstance(self.field.type, ast_nodes.AnonymousStructDecl)
and self.field.name not in self.manually_added_fields
and anonymous_struct_name in constants.STRUCTS_TO_BIND
):
return WrappedFieldData(
return structs_wrappers_data.WrappedFieldData(
binding=struct_field_code_builder.build_simple_property_binding(
self.field,
self.struct_wrapper_name,
@@ -294,24 +285,26 @@ class StructFieldHandler:
)
return self._get_manual_definition(comment_type="anonymous struct field")
def _undefined(self) -> WrappedFieldData:
def _undefined(self) -> structs_wrappers_data.WrappedFieldData:
"""This function adds a TODO comment for fields that are not handled by this class yet."""
return WrappedFieldData(
return structs_wrappers_data.WrappedFieldData(
definition=f"// TODO: UNDEFINED definition for {self.field.name}",
binding=f"// TODO: UNDEFINED binding for {self.field.name}",
)
def _get_manual_definition(self, comment_type: str = "") -> WrappedFieldData:
def _get_manual_definition(
self, comment_type: str = ""
) -> structs_wrappers_data.WrappedFieldData:
"""Helper method to generate a comment as a definition for manually added fields."""
if self.field.name in self.manually_added_fields:
return WrappedFieldData(
return structs_wrappers_data.WrappedFieldData(
definition=(
f"// {comment_type} is defined manually. {self.field.name}"
),
binding=self.simple_property_binding,
)
return WrappedFieldData(
return structs_wrappers_data.WrappedFieldData(
definition=(
f"// TODO: Define {comment_type} manually for {self.field.name}"
),