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:
committed by
Copybara-Service
parent
10130297c0
commit
c22d94e470
@@ -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}"
|
||||
),
|
||||
|
||||
Reference in New Issue
Block a user