More simplification to functions.py in WASM bindings

PiperOrigin-RevId: 834643175
Change-Id: Iff873f6f0916e0039ff6a44468afd9b4d8f98f37
This commit is contained in:
Matija Kecman
2025-11-20 01:05:26 -08:00
committed by Copybara-Service
parent 102f382c9d
commit a0d49ade1b
3 changed files with 62 additions and 142 deletions
+44 -80
View File
@@ -60,9 +60,12 @@ 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 get_pointer_return_inner_value_type(func) is not None or any(
get_inner_value_type(param) for param in func.parameters
)
if get_pointer_return_inner_value_type(func):
return True
for param in func.parameters:
if get_inner_value_type(param):
return True
return False
def generate_function_wrapper(func: ast_nodes.FunctionDecl) -> str:
@@ -88,11 +91,8 @@ def generate_function_wrapper(func: ast_nodes.FunctionDecl) -> str:
if bound_check_code:
builder.line(bound_check_code)
c_params_list = get_params_string_maybe_with_conversion(func.parameters)
c_params_str = ", ".join(c_params_list)
c_call = f"{func.name}({c_params_str})"
c_statement = get_compatible_return_call(func, c_call)
builder.line(f"{c_statement};")
for line in get_compatible_return_code(func).splitlines():
builder.line(line)
return builder.to_string()
@@ -194,7 +194,7 @@ def get_param_string(p: ast_nodes.FunctionParameterDecl) -> str:
def get_params_string_maybe_with_conversion(
ast_params: Tuple[ast_nodes.FunctionParameterDecl, ...],
) -> List[str]:
) -> str:
"""Generates C++ expressions for passing compatible params from JS to MuJoCo C-API functions."""
native_params = []
@@ -217,27 +217,27 @@ def get_params_string_maybe_with_conversion(
f"Unhandled parameter type for conversion: {p.type} for param"
f" '{p.name}'"
)
return native_params
return ", ".join(native_params)
def get_compatible_return_call(
func: ast_nodes.FunctionDecl, invoker: str
) -> str:
def get_compatible_return_code(func: ast_nodes.FunctionDecl) -> str:
"""Generates embind compatible return value conversion."""
c_params = get_params_string_maybe_with_conversion(func.parameters)
c_call = f"{func.name}({c_params})"
if isinstance(func.return_type, ast_nodes.ValueType):
if func.return_type.name == "void":
return invoker
return f"{c_call};"
if func.return_type.name in constants.PRIMITIVE_TYPES:
return f"return {invoker}"
return f"return {c_call};"
if inner_type := get_pointer_return_inner_value_type(func):
if inner_type.name == "char":
return f"return std::string({invoker})"
return f"return std::string({c_call});"
elif inner_type.name == "mjString":
return f"return *{invoker}"
return f"return *{c_call};"
elif inner_type.name not in constants.PRIMITIVE_TYPES:
return get_converted_struct_to_class(func, invoker)
return get_optional_return_code(func, c_call)
raise RuntimeError(
"Failed to calculate return value conversion for function"
@@ -262,21 +262,22 @@ def get_compatible_return_type(func: ast_nodes.FunctionDecl) -> str:
return "val"
def get_converted_struct_to_class(
func: ast_nodes.FunctionDecl, invoker: str
def get_optional_return_code(
func: ast_nodes.FunctionDecl, c_call: str
) -> str:
"""Generates a C++ function invocation for a struct return-type function."""
"""Generates code to return std::optional of the wrapped struct."""
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 = common.capitalize(struct_name)
return_str = f"{class_constructor}(result)"
return f"""{const_qualifier}{struct_name}* result = {invoker};
if (result == nullptr) {{
return std::nullopt;
}}
return {return_str}"""
builder = code_builder.CodeBuilder()
builder.line(f"{const_qualifier}{struct_name}* result = {c_call};")
with builder.block("if (result == nullptr)"):
builder.line("return std::nullopt;")
builder.line(f"return {common.capitalize(struct_name)}(result);")
return builder.to_string()
def is_excluded_function_name(func_name: str) -> bool:
@@ -291,63 +292,26 @@ class Generator:
"""Generates Embind bindings for MuJoCo functions."""
def __init__(self, functions: Mapping[str, ast_nodes.FunctionDecl]):
self.direct_bind_functions: List[ast_nodes.FunctionDecl] = []
self.wrapper_bind_functions: List[ast_nodes.FunctionDecl] = []
for func in functions.values():
if should_be_wrapped(func):
self.wrapper_bind_functions.append(func)
else:
self.direct_bind_functions.append(func)
def _generate_wrappers(self) -> str:
"""Generates Embind bindings for all functions that need wrappers."""
code = []
for func in self.wrapper_bind_functions:
if func.name not in constants.MANUAL_WRAPPER_FUNCTIONS:
wrapper_code = generate_function_wrapper(func)
code.append(wrapper_code)
return "\n\n".join(code)
def _generate_direct_bindable_functions(self) -> list[str]:
"""Generates Embind bindings for all directly bindable functions."""
result = []
for func in self.direct_bind_functions:
result.append(self._generate_function_binding(func))
return result
def _generate_function_binding(
self, func: ast_nodes.FunctionDecl, is_wrapper=False
) -> str:
"""Generates the Embind code for a single function."""
js_name, cpp_func = func.name, func.name
if is_wrapper:
cpp_func += "_wrapper"
return f'function("{js_name}", &{cpp_func});'
def _generate_wrapper_bindable_functions(self) -> list[str]:
"""Generates Embind bindings for all functions that need wrappers."""
result = []
for func in self.wrapper_bind_functions:
result.append(self._generate_function_binding(func, True))
return result
self.functions = functions
def generate(self) -> list[tuple[str, list[str]]]:
"""Generates the bindings file for all functions."""
wrapper_functions = self._generate_wrappers()
function_bindings = self._generate_direct_bindable_functions()
function_bindings += self._generate_wrapper_bindable_functions()
function_wrappers = []
for func in self.functions.values():
if should_be_wrapped(func):
if func.name not in constants.MANUAL_WRAPPER_FUNCTIONS:
function_wrappers.append(generate_function_wrapper(func))
function_bindings = []
for func in self.functions.values():
js_name = func.name
cpp_func = func.name
if should_be_wrapped(func):
cpp_func += "_wrapper"
function_bindings.append(f'function("{js_name}", &{cpp_func});')
return [
("// {{ WRAPPER_FUNCTIONS }}", [wrapper_functions]),
("// {{ WRAPPER_FUNCTIONS }}", ["\n\n".join(function_wrappers)]),
("// {{ FUNCTION_BINDINGS }}", ["\n".join(function_bindings)]),
]