Refactor WASM bindings in preparation for codegen directory structure update
PiperOrigin-RevId: 829459535 Change-Id: I4045e163a6e7fcee9d9585131d1e5e30f0b50b89
This commit is contained in:
committed by
Copybara-Service
parent
59debb50b1
commit
44220fcc51
@@ -15,36 +15,6 @@
|
||||
"""Utility functions for code generation."""
|
||||
|
||||
import os
|
||||
import pathlib
|
||||
|
||||
from wasm.codegen.helpers import constants
|
||||
|
||||
|
||||
def get_default_output_dir() -> str:
|
||||
"""Gets the default output directory (sibling of 'generated' folder)."""
|
||||
# Get the directory of the current file (generator/base.py)
|
||||
current_dir = pathlib.Path(__file__).parent
|
||||
# Go up one level to the project root and then down to 'generated'
|
||||
default_output_dir = str(current_dir.parent / "generated")
|
||||
return default_output_dir
|
||||
|
||||
|
||||
def get_file_path(
|
||||
template_dir: str, output_dir: str, filename: str
|
||||
) -> tuple[str, str]:
|
||||
"""Constructs the template and output file paths.
|
||||
|
||||
Args:
|
||||
template_dir: The directory containing the template files.
|
||||
output_dir: The directory where the generated files will be saved.
|
||||
filename: The name of the file.
|
||||
|
||||
Returns:
|
||||
A tuple containing the template file path and the output file path.
|
||||
"""
|
||||
template_file = f"wasm/codegen/{template_dir}/{filename}"
|
||||
output_file = f"wasm/codegen/{output_dir}/{filename}"
|
||||
return template_file, output_file
|
||||
|
||||
|
||||
def write_to_file(filepath: str, content: str) -> None:
|
||||
@@ -71,53 +41,26 @@ def uppercase_first_letter(input_string: str) -> str:
|
||||
return input_string[:1].upper() + input_string[1:]
|
||||
|
||||
|
||||
def try_cast_to_scalar_type(value: str) -> int | float | str:
|
||||
"""Tries to cast a string to an integer, then a float, otherwise returns the original string."""
|
||||
for type_ in [int, float]:
|
||||
try:
|
||||
return type_(value)
|
||||
except ValueError:
|
||||
continue
|
||||
return value
|
||||
|
||||
|
||||
def replace_lines_containing_marker(
|
||||
lines: list[str],
|
||||
marker_to_replace: str,
|
||||
replacement_content: str | list[str],
|
||||
marker: str,
|
||||
content: list[str],
|
||||
) -> list[str]:
|
||||
"""Replaces lines containing a specific marker with new content."""
|
||||
|
||||
new_lines = []
|
||||
replaced = False
|
||||
for line in lines:
|
||||
if not replaced and marker_to_replace in line:
|
||||
indentation = _get_indentation(line)
|
||||
if isinstance(replacement_content, str):
|
||||
new_lines.append(indentation + replacement_content)
|
||||
elif isinstance(replacement_content, list):
|
||||
for content_line in replacement_content:
|
||||
if not content_line.strip():
|
||||
continue
|
||||
indented_line = (
|
||||
indentation
|
||||
+ content_line.replace("\n", "\n" + indentation)
|
||||
+ "\n"
|
||||
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"
|
||||
)
|
||||
new_lines.append(indented_line)
|
||||
replaced = True
|
||||
else:
|
||||
new_lines.append(line)
|
||||
return new_lines
|
||||
|
||||
|
||||
def _get_indentation(line: str) -> str:
|
||||
"""Returns the indentation of the given line as a string of spaces."""
|
||||
|
||||
indentation = ""
|
||||
for char in line:
|
||||
if char == " ":
|
||||
indentation += " "
|
||||
else:
|
||||
break
|
||||
return indentation
|
||||
return lines[:i] + replacement_lines + lines[i + 1 :]
|
||||
return lines
|
||||
|
||||
@@ -64,11 +64,6 @@ class CommonUtilsTest(absltest.TestCase):
|
||||
common.uppercase_first_letter(" leading space"), " leading space"
|
||||
)
|
||||
|
||||
def test_try_cast_to_scalar_type(self):
|
||||
self.assertEqual(common.try_cast_to_scalar_type("123"), 123)
|
||||
self.assertEqual(common.try_cast_to_scalar_type("123.456"), 123.456)
|
||||
self.assertEqual(common.try_cast_to_scalar_type("abc"), "abc")
|
||||
|
||||
|
||||
class FunctionUtilsTest(absltest.TestCase):
|
||||
|
||||
|
||||
@@ -187,12 +187,6 @@ def _generate_field_data(
|
||||
binding=_simple_property_binding(f, w),
|
||||
is_primitive_or_fixed_size=True,
|
||||
)
|
||||
else:
|
||||
return WrappedFieldData(
|
||||
definition=f"// TODO: NOT IMPLEMENTED ARRAY wrapper for {f.name}",
|
||||
typename=_get_field_struct_type(f.type),
|
||||
binding=f"// TODO: NOT IMPLEMENTED ARRAY binding for {f.name}",
|
||||
)
|
||||
|
||||
elif isinstance(f.type, ast_nodes.PointerType):
|
||||
|
||||
@@ -280,10 +274,12 @@ def _generate_field_data(
|
||||
binding=_simple_property_binding(f, w),
|
||||
)
|
||||
|
||||
# SHOULD NOT OCCUR
|
||||
print("Error: field {f.name} not properly handled")
|
||||
return WrappedFieldData(
|
||||
definition=f"// TODO: UNDEFINED definition for {f.name}",
|
||||
definition=f"// Error: field {f.name} not properly handled.",
|
||||
typename=_get_field_struct_type(f.type),
|
||||
binding=f"// TODO: UNDEFINED binding for {f.name}",
|
||||
binding=f"// Error: field {f.name} not properly handled.",
|
||||
)
|
||||
|
||||
|
||||
@@ -512,6 +508,8 @@ def _build_struct_bindings(
|
||||
):
|
||||
"""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
|
||||
@@ -528,9 +526,8 @@ def _build_struct_bindings(
|
||||
builder.line(".constructor<const MjSpec &>()")
|
||||
elif w == "MjvScene":
|
||||
builder.line(".constructor<MjModel *, int>()")
|
||||
|
||||
is_mjs = w.startswith("Mjs")
|
||||
if not is_mjs and w not in ["MjData", "MjModel", "MjSpec"]:
|
||||
builder.line(".constructor<>()")
|
||||
elif not is_mjs:
|
||||
builder.line(".constructor<>()")
|
||||
|
||||
shallow_copy = use_shallow_copy(wrapped_fields)
|
||||
|
||||
Reference in New Issue
Block a user