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
@@ -12,8 +12,6 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from typing import TypeAlias
|
||||
|
||||
from absl.testing import absltest
|
||||
from introspect import ast_nodes
|
||||
|
||||
@@ -21,37 +19,37 @@ from wasm.codegen.helpers import constants
|
||||
from wasm.codegen.helpers import function_utils
|
||||
|
||||
|
||||
PrimitiveTypes: TypeAlias = constants.PRIMITIVE_TYPES
|
||||
ValueType: TypeAlias = ast_nodes.ValueType
|
||||
PointerType: TypeAlias = ast_nodes.PointerType
|
||||
ArrayType: TypeAlias = ast_nodes.ArrayType
|
||||
FunctionParameterDecl: TypeAlias = ast_nodes.FunctionParameterDecl
|
||||
FunctionDecl: TypeAlias = ast_nodes.FunctionDecl
|
||||
|
||||
|
||||
class FunctionUtilsTest(absltest.TestCase):
|
||||
|
||||
def setUp(self):
|
||||
super().setUp()
|
||||
self.struct_type = ValueType("MyStruct")
|
||||
self.ptr_to_int = PointerType(ValueType("int"))
|
||||
self.func_ret_ptr_int = FunctionDecl(
|
||||
"func_pi", PointerType(ValueType("int")), [], "doc"
|
||||
self.struct_type = ast_nodes.ValueType("MyStruct")
|
||||
self.ptr_to_int = ast_nodes.PointerType(ast_nodes.ValueType("int"))
|
||||
self.func_ret_ptr_int = ast_nodes.FunctionDecl(
|
||||
"func_pi", ast_nodes.PointerType(ast_nodes.ValueType("int")), [], "doc"
|
||||
)
|
||||
self.func_ret_ptr_struct = FunctionDecl(
|
||||
"func_ps", PointerType(ValueType("MyStruct")), [], "doc"
|
||||
self.func_ret_ptr_struct = ast_nodes.FunctionDecl(
|
||||
"func_ps",
|
||||
ast_nodes.PointerType(ast_nodes.ValueType("MyStruct")),
|
||||
[],
|
||||
"doc",
|
||||
)
|
||||
|
||||
def test_return_is_value_of_type(self):
|
||||
self.assertTrue(
|
||||
function_utils.return_is_value_of_type(
|
||||
FunctionDecl("func_i", ValueType("int"), [], "doc"), PrimitiveTypes
|
||||
ast_nodes.FunctionDecl(
|
||||
"func_i", ast_nodes.ValueType("int"), [], "doc"
|
||||
),
|
||||
constants.PRIMITIVE_TYPES,
|
||||
)
|
||||
)
|
||||
self.assertFalse(
|
||||
function_utils.return_is_value_of_type(
|
||||
FunctionDecl("func_s", ValueType("MyStruct"), [], "doc"),
|
||||
PrimitiveTypes,
|
||||
ast_nodes.FunctionDecl(
|
||||
"func_s", ast_nodes.ValueType("MyStruct"), [], "doc"
|
||||
),
|
||||
constants.PRIMITIVE_TYPES,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -72,21 +70,25 @@ class FunctionUtilsTest(absltest.TestCase):
|
||||
)
|
||||
|
||||
def test_param_is_primitive_value(self):
|
||||
param_prim_val = FunctionParameterDecl("prim_v", ValueType("int"))
|
||||
param_arr = FunctionParameterDecl(
|
||||
"arr_v", ArrayType(ValueType("int"), extents=(10,))
|
||||
param_prim_val = ast_nodes.FunctionParameterDecl(
|
||||
"prim_v", ast_nodes.ValueType("int")
|
||||
)
|
||||
param_arr = ast_nodes.FunctionParameterDecl(
|
||||
"arr_v", ast_nodes.ArrayType(ast_nodes.ValueType("int"), extents=(10,))
|
||||
)
|
||||
|
||||
self.assertTrue(function_utils.param_is_primitive_value(param_prim_val))
|
||||
self.assertFalse(function_utils.param_is_primitive_value(param_arr))
|
||||
|
||||
def test_param_is_pointer_to_primitive_value(self):
|
||||
param_ptr_to_prim = FunctionParameterDecl("p_prim", self.ptr_to_int)
|
||||
param_arr_of_prim = FunctionParameterDecl(
|
||||
"a_prim", ArrayType(ValueType("int"), extents=(10,))
|
||||
param_ptr_to_prim = ast_nodes.FunctionParameterDecl(
|
||||
"p_prim", self.ptr_to_int
|
||||
)
|
||||
param_ptr_to_struct = FunctionParameterDecl(
|
||||
name="p_struct", type=PointerType(inner_type=self.struct_type)
|
||||
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(
|
||||
function_utils.param_is_pointer_to_primitive_value(param_ptr_to_prim)
|
||||
@@ -99,14 +101,14 @@ class FunctionUtilsTest(absltest.TestCase):
|
||||
)
|
||||
|
||||
def test_param_is_pointer_to_struct(self):
|
||||
param_arr_of_struct = FunctionParameterDecl(
|
||||
"a_struct", ArrayType(self.struct_type, extents=(5,))
|
||||
param_arr_of_struct = ast_nodes.FunctionParameterDecl(
|
||||
"a_struct", ast_nodes.ArrayType(self.struct_type, extents=(5,))
|
||||
)
|
||||
param_ptr_to_struct = FunctionParameterDecl(
|
||||
"p_struct", PointerType(self.struct_type)
|
||||
param_ptr_to_struct = ast_nodes.FunctionParameterDecl(
|
||||
"p_struct", ast_nodes.PointerType(self.struct_type)
|
||||
)
|
||||
param_ptr_to_ptr = FunctionParameterDecl(
|
||||
"p_ptr", PointerType(self.ptr_to_int)
|
||||
param_ptr_to_ptr = ast_nodes.FunctionParameterDecl(
|
||||
"p_ptr", ast_nodes.PointerType(self.ptr_to_int)
|
||||
)
|
||||
self.assertTrue(
|
||||
function_utils.param_is_pointer_to_struct(param_arr_of_struct)
|
||||
@@ -119,43 +121,46 @@ class FunctionUtilsTest(absltest.TestCase):
|
||||
)
|
||||
|
||||
def test_should_be_wrapped_with_primitive_ptr_return(self):
|
||||
func = FunctionDecl(
|
||||
func = ast_nodes.FunctionDecl(
|
||||
name="get_data",
|
||||
return_type=PointerType(ValueType("int")),
|
||||
return_type=ast_nodes.PointerType(ast_nodes.ValueType("int")),
|
||||
parameters=tuple(),
|
||||
doc="Returns int pointer",
|
||||
)
|
||||
self.assertTrue(function_utils.should_be_wrapped(func))
|
||||
|
||||
def test_generate_function_wrapper_for_simple_func(self):
|
||||
func = FunctionDecl(
|
||||
func = ast_nodes.FunctionDecl(
|
||||
name="get_id",
|
||||
return_type=ValueType("int"),
|
||||
return_type=ast_nodes.ValueType("int"),
|
||||
parameters=tuple(),
|
||||
doc="Returns an integer ID",
|
||||
)
|
||||
result = function_utils.generate_function_wrapper(func)
|
||||
self.assertEqual(result, """int get_id_wrapper()
|
||||
self.assertEqual(
|
||||
result,
|
||||
"""int get_id_wrapper()
|
||||
{
|
||||
return get_id();
|
||||
}""")
|
||||
}""",
|
||||
)
|
||||
|
||||
def test_generate_function_wrapper_checking_param(self):
|
||||
parameters = (
|
||||
FunctionParameterDecl(
|
||||
ast_nodes.FunctionParameterDecl(
|
||||
name="mat",
|
||||
type=PointerType(
|
||||
inner_type=ValueType(name="mjtNum", is_const=True),
|
||||
type=ast_nodes.PointerType(
|
||||
inner_type=ast_nodes.ValueType(name="mjtNum", is_const=True),
|
||||
),
|
||||
),
|
||||
FunctionParameterDecl(
|
||||
ast_nodes.FunctionParameterDecl(
|
||||
name="nr",
|
||||
type=ValueType(name="int"),
|
||||
type=ast_nodes.ValueType(name="int"),
|
||||
),
|
||||
)
|
||||
func = FunctionDecl(
|
||||
func = ast_nodes.FunctionDecl(
|
||||
name="get_id",
|
||||
return_type=ValueType("int"),
|
||||
return_type=ast_nodes.ValueType("int"),
|
||||
parameters=parameters,
|
||||
doc="Returns an integer ID",
|
||||
)
|
||||
@@ -170,25 +175,25 @@ class FunctionUtilsTest(absltest.TestCase):
|
||||
)
|
||||
|
||||
def test_get_params_string_with_struct_ptr(self):
|
||||
param = FunctionParameterDecl(
|
||||
param = ast_nodes.FunctionParameterDecl(
|
||||
name="my_struct",
|
||||
type=PointerType(ValueType("mystruct")),
|
||||
type=ast_nodes.PointerType(ast_nodes.ValueType("mystruct")),
|
||||
)
|
||||
result = function_utils.get_params_string((param,))
|
||||
self.assertEqual(result, ["Mystruct& my_struct"])
|
||||
|
||||
def test_get_params_string_maybe_with_conversion_struct_ptr(self):
|
||||
param = FunctionParameterDecl(
|
||||
param = ast_nodes.FunctionParameterDecl(
|
||||
name="s",
|
||||
type=PointerType(ValueType("customstruct")),
|
||||
type=ast_nodes.PointerType(ast_nodes.ValueType("customstruct")),
|
||||
)
|
||||
result = function_utils.get_params_string_maybe_with_conversion((param,))
|
||||
self.assertEqual(result, ["s.get()"])
|
||||
|
||||
def test_get_compatible_return_call(self):
|
||||
func = FunctionDecl(
|
||||
func = ast_nodes.FunctionDecl(
|
||||
name="noop",
|
||||
return_type=ValueType("void"),
|
||||
return_type=ast_nodes.ValueType("void"),
|
||||
parameters=tuple(),
|
||||
doc="does nothing",
|
||||
)
|
||||
@@ -196,9 +201,9 @@ class FunctionUtilsTest(absltest.TestCase):
|
||||
self.assertEqual(result, "noop()")
|
||||
|
||||
def test_get_compatible_return_type(self):
|
||||
func = FunctionDecl(
|
||||
func = ast_nodes.FunctionDecl(
|
||||
name="get_name",
|
||||
return_type=PointerType(ValueType("char")),
|
||||
return_type=ast_nodes.PointerType(ast_nodes.ValueType("char")),
|
||||
parameters=tuple(),
|
||||
doc="returns name",
|
||||
)
|
||||
@@ -206,9 +211,9 @@ class FunctionUtilsTest(absltest.TestCase):
|
||||
self.assertEqual(result.strip(), "std::string")
|
||||
|
||||
def test_get_converted_struct_to_class(self):
|
||||
func = FunctionDecl(
|
||||
func = ast_nodes.FunctionDecl(
|
||||
name="get_struct",
|
||||
return_type=PointerType(ValueType("mystruct")),
|
||||
return_type=ast_nodes.PointerType(ast_nodes.ValueType("mystruct")),
|
||||
parameters=tuple(),
|
||||
doc="returns struct",
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user