Simplify functions.py in WASM bindings generation

PiperOrigin-RevId: 834470193
Change-Id: I3c26d65ea8812fd6a4c78264fbce8353bac6fb2f
This commit is contained in:
Matija Kecman
2025-11-19 15:34:34 -08:00
committed by Copybara-Service
parent a0cb97eb4d
commit 48ecddbe3b
2 changed files with 81 additions and 213 deletions
+81 -139
View File
@@ -23,76 +23,37 @@ 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 get_inner_value_type(
param: ast_nodes.FunctionParameterDecl,
) -> ast_nodes.ValueType | None:
if not isinstance(param.type, (ast_nodes.PointerType, ast_nodes.ArrayType)):
return None
if not isinstance(param.type.inner_type, ast_nodes.ValueType):
return None
return param.type.inner_type
def get_pointer_return_inner_value_type(
func: ast_nodes.FunctionDecl,
) -> ast_nodes.ValueType | None:
if not isinstance(func.return_type, ast_nodes.PointerType):
return None
if not isinstance(func.return_type.inner_type, ast_nodes.ValueType):
return None
return func.return_type.inner_type
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 param.type.name in constants.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
):
inner_type = get_pointer_return_inner_value_type(func)
if inner_type and inner_type.is_const:
return "const "
return ""
@@ -100,14 +61,8 @@ def get_const_qualifier(func: ast_nodes.FunctionDecl) -> str:
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
)
get_pointer_return_inner_value_type(func) is not None
or any(get_inner_value_type(param) for param in func.parameters)
)
@@ -164,30 +119,31 @@ def get_param_notnullable(
def get_param_unpack_statement(
p: ast_nodes.FunctionParameterDecl,
) -> str | None:
) -> str:
"""Generates C++ statements to unpack JS values for pointer/array parameters."""
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
return ""
if p.type.inner_type.is_const:
# param is Javascript number[]
if p.nullable:
return f"UNPACK_NULLABLE_ARRAY({p.type.inner_type.name}, {p.name});"
else:
return f"UNPACK_ARRAY({p.type.inner_type.name}, {p.name});"
inner_type = get_inner_value_type(p)
if not inner_type:
return ""
if inner_type.name not in constants.PRIMITIVE_TYPES:
return ""
if inner_type.name == "char":
# param is Javascript string
return ""
if inner_type.is_const:
# param is Javascript number[]
if p.nullable:
return f"UNPACK_NULLABLE_ARRAY({inner_type.name}, {p.name});"
else:
# param is TypedArray or a WasmBuffer
if p.nullable:
return f"UNPACK_NULLABLE_VALUE({p.type.inner_type.name}, {p.name});"
else:
return f"UNPACK_VALUE({p.type.inner_type.name}, {p.name});"
return None
return f"UNPACK_ARRAY({inner_type.name}, {p.name});"
else:
# param is TypedArray or a WasmBuffer
if p.nullable:
return f"UNPACK_NULLABLE_VALUE({inner_type.name}, {p.name});"
else:
return f"UNPACK_VALUE({inner_type.name}, {p.name});"
def get_param_string(
@@ -198,17 +154,17 @@ def get_param_string(
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
and p.type.inner_type.name not in constants.PRIMITIVE_TYPES
):
# Pointer to struct parameters
const_qualifier = "const " if p.type.inner_type.is_const else ""
return (
f"{const_qualifier}{uppercase_first_letter(p.type.inner_type.name)}&"
f"{const_qualifier}{common.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
and p.type.name in constants.PRIMITIVE_TYPES
):
# Primitive value parameters
const_qualifier = "const " if p.type.is_const else ""
@@ -216,7 +172,7 @@ def get_param_string(
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 in constants.PRIMITIVE_TYPES
):
# Pointer to primitive value parameters or arrays
if p.type.inner_type.name == "char":
@@ -248,26 +204,19 @@ def get_params_string_maybe_with_conversion(
native_params = []
for p in ast_params:
if param_is_pointer_to_struct(p):
native_params.append(f"{p.name}.get()")
if inner_type := get_inner_value_type(p):
if inner_type.name in constants.PRIMITIVE_TYPES:
if inner_type.name == "char":
const_qualifier = "const " if inner_type.is_const else ""
native_params.append(
f"{p.name}.as<{const_qualifier}std::string>().data()"
)
else:
native_params.append(f"{p.name}_.data()")
else: # struct
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"
@@ -281,19 +230,20 @@ def get_compatible_return_call(
) -> 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":
if isinstance(func.return_type, ast_nodes.ValueType):
if func.return_type.name == "void":
return invoker
if func.return_type.name in constants.PRIMITIVE_TYPES:
return f"return {invoker}"
if inner_type := get_pointer_return_inner_value_type(func):
if inner_type.name == "char":
return f"return std::string({invoker})"
elif func.return_type.inner_type.name == "mjString":
elif 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}"
elif inner_type.name not in constants.PRIMITIVE_TYPES:
return get_converted_struct_to_class(func, invoker)
raise RuntimeError(
"Failed to calculate return value conversion for function"
f" {func.name} that returns '{func.return_type}'"
@@ -303,22 +253,15 @@ def get_compatible_return_call(
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 inner_type := get_pointer_return_inner_value_type(func):
if inner_type.name in ["char", "mjString"]:
return "std::string"
if inner_type.name not in constants.PRIMITIVE_TYPES:
const_qualifier = get_const_qualifier(func)
return f"""{const_qualifier}std::optional<{common.uppercase_first_letter(inner_type.name)}>"""
if (
isinstance(func.return_type, ast_nodes.ValueType)
and func.return_type.name in PRIMITIVE_TYPES
and func.return_type.name in constants.PRIMITIVE_TYPES
):
return f"{func.return_type.name}"
return "val"
@@ -332,7 +275,7 @@ def get_converted_struct_to_class(
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)
class_constructor = common.uppercase_first_letter(struct_name)
return_str = f"{class_constructor}(result)"
return f"""{const_qualifier}{struct_name}* result = {invoker};
if (result == nullptr) {{
@@ -344,8 +287,7 @@ def get_converted_struct_to_class(
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_")
func_name.startswith(("mjr_", "mjui_"))
or func_name in constants.SKIPPED_FUNCTIONS
)
-74
View File
@@ -82,40 +82,6 @@ class FunctionUtilsTest(absltest.TestCase):
"doc",
)
def test_return_is_value_of_type(self):
self.assertTrue(
functions.return_is_value_of_type(
ast_nodes.FunctionDecl(
"func_i", ast_nodes.ValueType("int"), [], "doc"
),
constants.PRIMITIVE_TYPES,
)
)
self.assertFalse(
functions.return_is_value_of_type(
ast_nodes.FunctionDecl(
"func_s", ast_nodes.ValueType("MyStruct"), [], "doc"
),
constants.PRIMITIVE_TYPES,
)
)
def test_return_is_pointer_to_struct(self):
self.assertTrue(
functions.return_is_pointer_to_struct(self.func_ret_ptr_struct)
)
self.assertFalse(
functions.return_is_pointer_to_struct(self.func_ret_ptr_int)
)
def test_return_is_pointer_to_primitive(self):
self.assertTrue(
functions.return_is_pointer_to_primitive(self.func_ret_ptr_int)
)
self.assertFalse(
functions.return_is_pointer_to_primitive(self.func_ret_ptr_struct)
)
def test_param_is_primitive_value(self):
param_prim_val = ast_nodes.FunctionParameterDecl(
"prim_v", ast_nodes.ValueType("int")
@@ -127,46 +93,6 @@ class FunctionUtilsTest(absltest.TestCase):
self.assertTrue(functions.param_is_primitive_value(param_prim_val))
self.assertFalse(functions.param_is_primitive_value(param_arr))
def test_param_is_pointer_to_primitive_value(self):
param_ptr_to_prim = ast_nodes.FunctionParameterDecl(
"p_prim", self.ptr_to_int
)
param_arr_of_prim = ast_nodes.FunctionParameterDecl(
"a_prim", ast_nodes.ArrayType(ast_nodes.ValueType("int"), extents=(10,))
)
param_ptr_to_struct = ast_nodes.FunctionParameterDecl(
name="p_struct", type=ast_nodes.PointerType(inner_type=self.struct_type)
)
self.assertTrue(
functions.param_is_pointer_to_primitive_value(param_ptr_to_prim)
)
self.assertTrue(
functions.param_is_pointer_to_primitive_value(param_arr_of_prim)
)
self.assertFalse(
functions.param_is_pointer_to_primitive_value(param_ptr_to_struct)
)
def test_param_is_pointer_to_struct(self):
param_arr_of_struct = ast_nodes.FunctionParameterDecl(
"a_struct", ast_nodes.ArrayType(self.struct_type, extents=(5,))
)
param_ptr_to_struct = ast_nodes.FunctionParameterDecl(
"p_struct", ast_nodes.PointerType(self.struct_type)
)
param_ptr_to_ptr = ast_nodes.FunctionParameterDecl(
"p_ptr", ast_nodes.PointerType(self.ptr_to_int)
)
self.assertTrue(
functions.param_is_pointer_to_struct(param_arr_of_struct)
)
self.assertTrue(
functions.param_is_pointer_to_struct(param_ptr_to_struct)
)
self.assertFalse(
functions.param_is_pointer_to_struct(param_ptr_to_ptr)
)
def test_should_be_wrapped_with_primitive_ptr_return(self):
func = ast_nodes.FunctionDecl(
name="get_data",