From 48ecddbe3b1c6b9d7a0df96373e3f733603c8194 Mon Sep 17 00:00:00 2001 From: Matija Kecman Date: Wed, 19 Nov 2025 15:34:34 -0800 Subject: [PATCH] Simplify functions.py in WASM bindings generation PiperOrigin-RevId: 834470193 Change-Id: I3c26d65ea8812fd6a4c78264fbce8353bac6fb2f --- wasm/codegen/generators/functions.py | 220 ++++++++++---------------- wasm/codegen/tests/generators_test.py | 74 --------- 2 files changed, 81 insertions(+), 213 deletions(-) diff --git a/wasm/codegen/generators/functions.py b/wasm/codegen/generators/functions.py index 7a2d4144..fc4d49af 100644 --- a/wasm/codegen/generators/functions.py +++ b/wasm/codegen/generators/functions.py @@ -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 ) diff --git a/wasm/codegen/tests/generators_test.py b/wasm/codegen/tests/generators_test.py index 272d85ea..55a42132 100644 --- a/wasm/codegen/tests/generators_test.py +++ b/wasm/codegen/tests/generators_test.py @@ -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",