Use a dedicated function for wrapper struct names

This change introduces a `wrapper_struct_name` function in `structs.py` to consistently generate the name of the C++ wrapper struct from a C struct name. All call sites in `functions.py` and `structs.py` that previously used `common.capitalize` for this purpose now call `wrapper_struct_name`.

PiperOrigin-RevId: 837544242
Change-Id: Ibe5fea8294283dae63bb8d064fd8beb8cfd3f2fe
This commit is contained in:
Google DeepMind
2025-11-27 08:41:25 -08:00
committed by Copybara-Service
parent 9ca1598b23
commit 8a5c52d395
11 changed files with 146 additions and 98 deletions
+2 -2
View File
@@ -7803,7 +7803,7 @@ MjSpec::~MjSpec() {
mjSpec *MjSpec::get() const { return ptr_; } mjSpec *MjSpec::get() const { return ptr_; }
void MjSpec::set(mjSpec *ptr) { ptr_ = ptr; } void MjSpec::set(mjSpec *ptr) { ptr_ = ptr; }
std::unique_ptr<MjModel> loadFromXML_wrapper(std::string filename) { std::unique_ptr<MjModel> mj_loadXML_wrapper(std::string filename) {
char error[1000]; char error[1000];
mjModel *model = mj_loadXML(filename.c_str(), nullptr, error, sizeof(error)); mjModel *model = mj_loadXML(filename.c_str(), nullptr, error, sizeof(error));
if (!model) { if (!model) {
@@ -10940,7 +10940,7 @@ EMSCRIPTEN_BINDINGS(mujoco_bindings) {
.property("useexisting", &MjLROpt::useexisting, &MjLROpt::set_useexisting, reference()) .property("useexisting", &MjLROpt::useexisting, &MjLROpt::set_useexisting, reference())
.property("uselimit", &MjLROpt::uselimit, &MjLROpt::set_uselimit, reference()); .property("uselimit", &MjLROpt::uselimit, &MjLROpt::set_uselimit, reference());
emscripten::class_<MjModel>("MjModel") emscripten::class_<MjModel>("MjModel")
.class_function("loadFromXML", &loadFromXML_wrapper, take_ownership()) .class_function("mj_loadXML", &mj_loadXML_wrapper, take_ownership())
.constructor<const MjModel &>() .constructor<const MjModel &>()
.property("B_colind", &MjModel::B_colind) .property("B_colind", &MjModel::B_colind)
.property("B_rowadr", &MjModel::B_rowadr) .property("B_rowadr", &MjModel::B_rowadr)
@@ -12,7 +12,6 @@
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
from introspect import ast_nodes
from introspect import enums as introspect_enums from introspect import enums as introspect_enums
from introspect import functions as introspect_functions from introspect import functions as introspect_functions
+50
View File
@@ -15,6 +15,7 @@
"""Utility functions for code generation.""" """Utility functions for code generation."""
import os import os
from introspect import ast_nodes
def write_to_file(filepath: str, content: str) -> None: def write_to_file(filepath: str, content: str) -> None:
@@ -64,3 +65,52 @@ def replace_lines_containing_marker(
) )
return lines[:i] + replacement_lines + lines[i + 1 :] return lines[:i] + replacement_lines + lines[i + 1 :]
return lines return lines
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 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 should_be_wrapped(func: ast_nodes.FunctionDecl) -> bool:
"""Checks if a MuJoCo function needs a wrapper function."""
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 wrapped_struct_name(c_struct_name: str) -> str:
"""Returns the name of the struct wrapping the given C struct."""
return capitalize(c_struct_name)
def wrapped_function_name(func: ast_nodes.FunctionDecl) -> str:
"""Returns the name of the function wrapping the given C function.
Hard-coded wrappers in the template need changing if the implementation of the
wrapped name is changed
Args:
func: The FunctionDecl of function to wrap.
"""
if should_be_wrapped(func):
return f"{func.name}_wrapper"
return func.name
+4 -4
View File
@@ -14,7 +14,7 @@
"""Constants used in the code generation process.""" """Constants used in the code generation process."""
from typing import Dict, List, Set from typing import Dict, Set
from introspect import structs as introspect_structs from introspect import structs as introspect_structs
PRIMITIVE_TYPES: Set[str] = { PRIMITIVE_TYPES: Set[str] = {
@@ -263,7 +263,7 @@ ANONYMOUS_STRUCTS: Dict[str, Dict[str, str]] = {
# This list is created by subtracting the skipped structs from the list of all # This list is created by subtracting the skipped structs from the list of all
# structs and adding the anonymous structs. # structs and adding the anonymous structs.
STRUCTS_TO_BIND: List[str] = list( STRUCTS_TO_BIND: list[str] = list(
set(introspect_structs.STRUCTS.keys()) set(introspect_structs.STRUCTS.keys())
.union(ANONYMOUS_STRUCTS.keys()) .union(ANONYMOUS_STRUCTS.keys())
.difference(set(SKIPPED_STRUCTS)) .difference(set(SKIPPED_STRUCTS))
@@ -363,10 +363,10 @@ MJDATA_SIZES = (
) )
# Fields that should be entirely omitted from the bindings. # Fields that should be entirely omitted from the bindings.
SKIPPED_FIELDS: Dict[str, List[str]] = {} SKIPPED_FIELDS: Dict[str, list[str]] = {}
# Fields handled manually in template file struct declaration. # Fields handled manually in template file struct declaration.
MANUAL_FIELDS: Dict[str, List[str]] = { MANUAL_FIELDS: Dict[str, list[str]] = {
# go/keep-sorted start # go/keep-sorted start
"MjData": ["contact"], "MjData": ["contact"],
"MjvScene": [ "MjvScene": [
+1 -3
View File
@@ -14,15 +14,13 @@
"""Generates Embind bindings for MuJoCo enums.""" """Generates Embind bindings for MuJoCo enums."""
from typing import List
from introspect import ast_nodes from introspect import ast_nodes
from wasm.codegen.generators import code_builder from wasm.codegen.generators import code_builder
def generate( def generate(
enums: List[ast_nodes.EnumDecl], enums: list[ast_nodes.EnumDecl],
) -> list[tuple[str, list[str]]]: ) -> list[tuple[str, list[str]]]:
"""Generates all Embind code for the provided enums.""" """Generates all Embind code for the provided enums."""
+13 -42
View File
@@ -14,7 +14,7 @@
"""Helper functions for processing and generating bindings for MuJoCo functions.""" """Helper functions for processing and generating bindings for MuJoCo functions."""
from typing import List, Tuple, cast from typing import Tuple, cast
from introspect import ast_nodes from introspect import ast_nodes
@@ -23,26 +23,6 @@ from wasm.codegen.generators import common
from wasm.codegen.generators import constants from wasm.codegen.generators import constants
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: def param_is_primitive_value(param: ast_nodes.FunctionParameterDecl) -> bool:
"""Checks if param is a primitive value type.""" """Checks if param is a primitive value type."""
if isinstance(param.type, ast_nodes.ValueType): if isinstance(param.type, ast_nodes.ValueType):
@@ -52,22 +32,12 @@ def param_is_primitive_value(param: ast_nodes.FunctionParameterDecl) -> bool:
def get_const_qualifier(func: ast_nodes.FunctionDecl) -> str: def get_const_qualifier(func: ast_nodes.FunctionDecl) -> str:
"""Returns the const qualifier of func's return type.""" """Returns the const qualifier of func's return type."""
inner_type = get_pointer_return_inner_value_type(func) inner_type = common.get_pointer_return_inner_value_type(func)
if inner_type and inner_type.is_const: if inner_type and inner_type.is_const:
return "const " return "const "
return "" return ""
def should_be_wrapped(func: ast_nodes.FunctionDecl) -> bool:
"""Checks if a MuJoCo function needs a wrapper function."""
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: def generate_function_wrapper(func: ast_nodes.FunctionDecl) -> str:
"""Generates C++ code for a wrapper function.""" """Generates C++ code for a wrapper function."""
@@ -78,7 +48,8 @@ def generate_function_wrapper(func: ast_nodes.FunctionDecl) -> str:
wrapper_params = ", ".join([get_param_string(p) for p in wrapper_parameters]) wrapper_params = ", ".join([get_param_string(p) for p in wrapper_parameters])
ret_type = get_compatible_return_type(func) ret_type = get_compatible_return_type(func)
builder = code_builder.CodeBuilder() builder = code_builder.CodeBuilder()
with builder.function(f"{ret_type} {func.name}_wrapper({wrapper_params})"): w = common.wrapped_function_name(func)
with builder.function(f"{ret_type} {w}({wrapper_params})"):
for p in wrapper_parameters: for p in wrapper_parameters:
if c_notnullable := get_param_notnullable(p): if c_notnullable := get_param_notnullable(p):
@@ -119,7 +90,7 @@ def get_param_unpack_statement(
) -> str: ) -> str:
"""Generates C++ statements to unpack JS values for pointer/array parameters.""" """Generates C++ statements to unpack JS values for pointer/array parameters."""
inner_type = get_inner_value_type(p) inner_type = common.get_inner_value_type(p)
if not inner_type: if not inner_type:
return "" return ""
@@ -199,7 +170,7 @@ def get_params_string_maybe_with_conversion(
native_params = [] native_params = []
for p in ast_params: for p in ast_params:
if inner_type := get_inner_value_type(p): if inner_type := common.get_inner_value_type(p):
if inner_type.name in constants.PRIMITIVE_TYPES: if inner_type.name in constants.PRIMITIVE_TYPES:
if inner_type.name == "char": if inner_type.name == "char":
const_qualifier = "const " if inner_type.is_const else "" const_qualifier = "const " if inner_type.is_const else ""
@@ -231,7 +202,7 @@ def get_compatible_return_code(func: ast_nodes.FunctionDecl) -> str:
if func.return_type.name in constants.PRIMITIVE_TYPES: if func.return_type.name in constants.PRIMITIVE_TYPES:
return f"return {c_call};" return f"return {c_call};"
if inner_type := get_pointer_return_inner_value_type(func): if inner_type := common.get_pointer_return_inner_value_type(func):
if inner_type.name == "char": if inner_type.name == "char":
return f"return std::string({c_call});" return f"return std::string({c_call});"
elif inner_type.name == "mjString": elif inner_type.name == "mjString":
@@ -248,7 +219,7 @@ def get_compatible_return_code(func: ast_nodes.FunctionDecl) -> str:
def get_compatible_return_type(func: ast_nodes.FunctionDecl) -> str: def get_compatible_return_type(func: ast_nodes.FunctionDecl) -> str:
"""Creates embind compatible return type.""" """Creates embind compatible return type."""
if inner_type := get_pointer_return_inner_value_type(func): if inner_type := common.get_pointer_return_inner_value_type(func):
if inner_type.name in ["char", "mjString"]: if inner_type.name in ["char", "mjString"]:
return "std::string" return "std::string"
if inner_type.name not in constants.PRIMITIVE_TYPES: if inner_type.name not in constants.PRIMITIVE_TYPES:
@@ -275,7 +246,7 @@ def get_optional_return_code(
builder.line(f"{const_qualifier}{struct_name}* result = {c_call};") builder.line(f"{const_qualifier}{struct_name}* result = {c_call};")
with builder.block("if (result == nullptr)"): with builder.block("if (result == nullptr)"):
builder.line("return std::nullopt;") builder.line("return std::nullopt;")
builder.line(f"return {common.capitalize(struct_name)}(result);") builder.line(f"return {common.wrapped_struct_name(struct_name)}(result);")
return builder.to_string() return builder.to_string()
@@ -289,20 +260,20 @@ def is_excluded_function_name(func_name: str) -> bool:
def generate( def generate(
functions: List[ast_nodes.FunctionDecl], functions: list[ast_nodes.FunctionDecl],
) -> list[tuple[str, list[str]]]: ) -> list[tuple[str, list[str]]]:
"""Generates Embind bindings for MuJoCo functions.""" """Generates Embind bindings for MuJoCo functions."""
wrapper_functions = [] wrapper_functions = []
for func in sorted(functions, key=lambda f: f.name): for func in sorted(functions, key=lambda f: f.name):
if should_be_wrapped(func): if common.should_be_wrapped(func):
if func.name not in constants.MANUAL_WRAPPER_FUNCTIONS: if func.name not in constants.MANUAL_WRAPPER_FUNCTIONS:
wrapper_functions.append(generate_function_wrapper(func)) wrapper_functions.append(generate_function_wrapper(func))
wrapper_content = "\n\n".join(wrapper_functions) wrapper_content = "\n\n".join(wrapper_functions)
function_bindings = [] function_bindings = []
for func in sorted(functions, key=lambda f: f.name): for func in sorted(functions, key=lambda f: f.name):
suffix = "_wrapper" if should_be_wrapped(func) else "" w = common.wrapped_function_name(func)
function_bindings.append(f'function("{func.name}", &{func.name}{suffix});') function_bindings.append(f'function("{func.name}", &{w});')
bindings_content = "\n".join(function_bindings) bindings_content = "\n".join(function_bindings)
return [ return [
+14 -11
View File
@@ -21,6 +21,7 @@ from typing import Tuple, Union, cast
from introspect import ast_nodes from introspect import ast_nodes
from introspect import structs as introspect_structs from introspect import structs as introspect_structs
from introspect import functions as introspect_functions
from wasm.codegen.generators import code_builder from wasm.codegen.generators import code_builder
from wasm.codegen.generators import common from wasm.codegen.generators import common
@@ -123,7 +124,9 @@ def _generate_field_data(
return WrappedFieldData( return WrappedFieldData(
binding=_get_property_binding(f, w, setter=False, reference=True), binding=_get_property_binding(f, w, setter=False, reference=True),
typename=anonymous_struct_name, typename=anonymous_struct_name,
declaration=f"{common.capitalize(anonymous_struct_name)} {f.name};", declaration=(
f"{common.wrapped_struct_name(anonymous_struct_name)} {f.name};"
),
ptr_initialization=f"{f.name}(&ptr_->{f.name})", ptr_initialization=f"{f.name}(&ptr_->{f.name})",
ptr_copy_reset=f"{f.name}.set(&ptr_->{f.name});", ptr_copy_reset=f"{f.name}.set(&ptr_->{f.name});",
is_primitive_or_fixed_size=True, is_primitive_or_fixed_size=True,
@@ -320,7 +323,7 @@ def build_struct_header(
): ):
"""Builds the C++ header file code for a struct.""" """Builds the C++ header file code for a struct."""
s = struct_name s = struct_name
w = common.capitalize(s) w = common.wrapped_struct_name(s)
if ( if (
s not in constants.ANONYMOUS_STRUCTS s not in constants.ANONYMOUS_STRUCTS
@@ -398,7 +401,7 @@ def build_struct_source(
): ):
"""Builds the C++ .cc file code for a struct.""" """Builds the C++ .cc file code for a struct."""
s = struct_name s = struct_name
w = common.capitalize(s) w = common.wrapped_struct_name(s)
is_mjs = w.startswith("Mjs") is_mjs = w.startswith("Mjs")
member_inits = _find_member_inits(wrapped_fields) member_inits = _find_member_inits(wrapped_fields)
@@ -467,7 +470,7 @@ def _build_struct_bindings(
wrapped_fields: list[WrappedFieldData], wrapped_fields: list[WrappedFieldData],
): ):
"""Builds the C++ bindings for a struct.""" """Builds the C++ bindings for a struct."""
w = common.capitalize(struct_name) w = common.wrapped_struct_name(struct_name)
is_mjs = w.startswith("Mjs") is_mjs = w.startswith("Mjs")
builder = code_builder.CodeBuilder() builder = code_builder.CodeBuilder()
@@ -478,10 +481,10 @@ def _build_struct_bindings(
builder.line(".constructor<MjModel *>()") builder.line(".constructor<MjModel *>()")
builder.line(".constructor<const MjModel &, const MjData &>()") builder.line(".constructor<const MjModel &, const MjData &>()")
elif w == "MjModel": elif w == "MjModel":
builder.line( w = common.wrapped_function_name(
'.class_function("loadFromXML", &loadFromXML_wrapper,' introspect_functions.FUNCTIONS["mj_loadXML"]
" take_ownership())"
) )
builder.line(f'.class_function("mj_loadXML", &{w}, take_ownership())')
builder.line(".constructor<const MjModel &>()") builder.line(".constructor<const MjModel &>()")
elif w == "MjSpec": elif w == "MjSpec":
builder.line(".constructor<const MjSpec &>()") builder.line(".constructor<const MjSpec &>()")
@@ -572,7 +575,7 @@ def _get_field_struct_type(
) -> str | None: ) -> str | None:
"""Extracts the base struct name if the field type is a struct or pointer to a struct.""" """Extracts the base struct name if the field type is a struct or pointer to a struct."""
s = struct_name s = struct_name
w = common.capitalize(s) w = common.wrapped_struct_name(s)
if isinstance(field.type, ast_nodes.AnonymousStructDecl): if isinstance(field.type, ast_nodes.AnonymousStructDecl):
anonymous_struct_name = "" anonymous_struct_name = ""
for name, value in constants.ANONYMOUS_STRUCTS.items(): for name, value in constants.ANONYMOUS_STRUCTS.items():
@@ -692,7 +695,7 @@ def generate(struct_to_bind: list[str]) -> list[tuple[str, list[str]]]:
fields: list[WrappedFieldData] = [] fields: list[WrappedFieldData] = []
introspect_fields = get_introspect_struct_fields(s) introspect_fields = get_introspect_struct_fields(s)
for field in introspect_fields: for field in introspect_fields:
fields.append(_generate_field_data(field, common.capitalize(s))) fields.append(_generate_field_data(field, common.wrapped_struct_name(s)))
wrapped_structs_with_fields[s] = fields wrapped_structs_with_fields[s] = fields
dependency_sorted_struct_names = sort_structs_by_dependency( dependency_sorted_struct_names = sort_structs_by_dependency(
@@ -723,13 +726,13 @@ def generate(struct_to_bind: list[str]) -> list[tuple[str, list[str]]]:
bindings.append(_build_struct_bindings(s, fields)) bindings.append(_build_struct_bindings(s, fields))
for s in alphabetically_sorted_struct_names: for s in alphabetically_sorted_struct_names:
w = common.capitalize(s) w = common.wrapped_struct_name(s)
if w.startswith("Mjs") or w == "MjSpec": if w.startswith("Mjs") or w == "MjSpec":
bindings.append(f"emscripten::register_optional<{w}>();") bindings.append(f"emscripten::register_optional<{w}>();")
manual_struct_field_declarations = [] manual_struct_field_declarations = []
for s in dependency_sorted_struct_names: for s in dependency_sorted_struct_names:
w = common.capitalize(s) w = common.wrapped_struct_name(s)
fields = wrapped_structs_with_fields[s] fields = wrapped_structs_with_fields[s]
if s in constants.MANUAL_STRUCTS_HEADERS: if s in constants.MANUAL_STRUCTS_HEADERS:
decls: list[str] = [] decls: list[str] = []
+1 -1
View File
@@ -386,7 +386,7 @@ MjSpec::~MjSpec() {
mjSpec *MjSpec::get() const { return ptr_; } mjSpec *MjSpec::get() const { return ptr_; }
void MjSpec::set(mjSpec *ptr) { ptr_ = ptr; } void MjSpec::set(mjSpec *ptr) { ptr_ = ptr; }
std::unique_ptr<MjModel> loadFromXML_wrapper(std::string filename) { std::unique_ptr<MjModel> mj_loadXML_wrapper(std::string filename) {
char error[1000]; char error[1000];
mjModel *model = mj_loadXML(filename.c_str(), nullptr, error, sizeof(error)); mjModel *model = mj_loadXML(filename.c_str(), nullptr, error, sizeof(error));
if (!model) { if (!model) {
+52 -25
View File
@@ -97,48 +97,75 @@ class FunctionUtilsTest(absltest.TestCase):
parameters=tuple(), parameters=tuple(),
doc="Returns int pointer", doc="Returns int pointer",
) )
self.assertTrue(functions.should_be_wrapped(func)) self.assertTrue(common.should_be_wrapped(func))
def test_generate_function_wrapper_for_simple_func(self): def test_generate_function_wrapper_for_simple_func(self):
func = ast_nodes.FunctionDecl( func = ast_nodes.FunctionDecl(
name="get_id", name="mj_defaultLROpt",
return_type=ast_nodes.ValueType("int"), return_type=ast_nodes.ValueType("void"),
parameters=tuple(), parameters=(
doc="Returns an integer ID", ast_nodes.FunctionParameterDecl(
name="opt",
type=ast_nodes.PointerType(
inner_type=ast_nodes.ValueType("mjLROpt"),
),
),
),
doc="Set default options for length range computation.",
) )
result = functions.generate_function_wrapper(func) result = functions.generate_function_wrapper(func)
self.assertEqual( self.assertEqual(
result, result,
"""int get_id_wrapper() { """void mj_defaultLROpt_wrapper(MjLROpt& opt) {
return get_id(); mj_defaultLROpt(opt.get());
}""", }""",
) )
def test_generate_function_wrapper_checking_param(self): def test_generate_function_wrapper_checking_param(self):
parameters = ( func = ast_nodes.FunctionDecl(
ast_nodes.FunctionParameterDecl( name="mj_extractState",
name="mat", return_type=ast_nodes.ValueType(name="void"),
type=ast_nodes.PointerType( parameters=(
inner_type=ast_nodes.ValueType(name="mjtNum", is_const=True), ast_nodes.FunctionParameterDecl(
name="m",
type=ast_nodes.PointerType(
inner_type=ast_nodes.ValueType(
name="mjModel", is_const=True
),
),
),
ast_nodes.FunctionParameterDecl(
name="src",
type=ast_nodes.PointerType(
inner_type=ast_nodes.ValueType(
name="mjtNum", is_const=True
),
),
),
ast_nodes.FunctionParameterDecl(
name="srcsig",
type=ast_nodes.ValueType(name="unsigned int"),
),
ast_nodes.FunctionParameterDecl(
name="dst",
type=ast_nodes.PointerType(
inner_type=ast_nodes.ValueType(name="mjtNum"),
),
),
ast_nodes.FunctionParameterDecl(
name="dstsig",
type=ast_nodes.ValueType(name="unsigned int"),
), ),
), ),
ast_nodes.FunctionParameterDecl( doc="Extract a subset of components from a state previously obtained via mj_getState.", # pylint: disable=line-too-long
name="nr",
type=ast_nodes.ValueType(name="int"),
),
)
func = ast_nodes.FunctionDecl(
name="get_id",
return_type=ast_nodes.ValueType("int"),
parameters=parameters,
doc="Returns an integer ID",
) )
result = functions.generate_function_wrapper(func) result = functions.generate_function_wrapper(func)
self.assertEqual( self.assertEqual(
result, result,
"""int get_id_wrapper(const NumberArray& mat, int nr) { """void mj_extractState_wrapper(const MjModel& m, const NumberArray& src, unsigned int srcsig, const val& dst, unsigned int dstsig) {
UNPACK_ARRAY(mjtNum, mat); UNPACK_ARRAY(mjtNum, src);
return get_id(mat_.data(), nr); UNPACK_VALUE(mjtNum, dst);
mj_extractState(m.get(), src_.data(), srcsig, dst_.data(), dstsig);
}""", }""",
) )
+8 -8
View File
@@ -119,7 +119,7 @@ describe('MuJoCo WASM Bindings', () => {
writeXMLFile(tempXmlFilename, TEST_XML); writeXMLFile(tempXmlFilename, TEST_XML);
model = mujoco.MjModel!.loadFromXML(tempXmlFilename); model = mujoco.MjModel!.mj_loadXML(tempXmlFilename);
if (!model) { if (!model) {
unlinkXMLFile(tempXmlFilename); unlinkXMLFile(tempXmlFilename);
throw new Error('Failed to load model from XML'); throw new Error('Failed to load model from XML');
@@ -182,7 +182,7 @@ describe('MuJoCo WASM Bindings', () => {
beforeEach(() => { beforeEach(() => {
writeXMLFile(tempXmlFilename, simpleXmlContent); writeXMLFile(tempXmlFilename, simpleXmlContent);
simpleModel = mujoco.MjModel!.loadFromXML(tempXmlFilename); simpleModel = mujoco.MjModel!.mj_loadXML(tempXmlFilename);
assertExists(simpleModel); assertExists(simpleModel);
simpleData = new mujoco.MjData(simpleModel); simpleData = new mujoco.MjData(simpleModel);
assertExists(simpleData); assertExists(simpleData);
@@ -345,7 +345,7 @@ describe('MuJoCo WASM Bindings', () => {
`; `;
const tempXmlFilename = '/tmp/model_c.xml'; const tempXmlFilename = '/tmp/model_c.xml';
writeXMLFile(tempXmlFilename, xmlString); writeXMLFile(tempXmlFilename, xmlString);
const model = mujoco.MjModel!.loadFromXML(tempXmlFilename); const model = mujoco.MjModel!.mj_loadXML(tempXmlFilename);
expect(model).not.toBeNull(); expect(model).not.toBeNull();
const data = new mujoco.MjData(model!); const data = new mujoco.MjData(model!);
expect(data).not.toBeNull(); expect(data).not.toBeNull();
@@ -879,7 +879,7 @@ describe('MuJoCo WASM Bindings', () => {
writeXMLFile(model1XmlFilename, xml2); writeXMLFile(model1XmlFilename, xml2);
writeXMLFile(model2XmlFilename, xml3); writeXMLFile(model2XmlFilename, xml3);
const model = mujoco.MjModel!.loadFromXML(modelXmlFilename); const model = mujoco.MjModel!.mj_loadXML(modelXmlFilename);
try { try {
expect(model).toBeDefined(); expect(model).toBeDefined();
@@ -1078,7 +1078,7 @@ describe('MuJoCo WASM Bindings', () => {
it('should check MjContactVec equality', () => { it('should check MjContactVec equality', () => {
const tempXmlFilename2 = '/tmp/model2.xml'; const tempXmlFilename2 = '/tmp/model2.xml';
writeXMLFile(tempXmlFilename2, TEST_XML); writeXMLFile(tempXmlFilename2, TEST_XML);
const model2 = mujoco.MjModel!.loadFromXML(tempXmlFilename2); const model2 = mujoco.MjModel!.mj_loadXML(tempXmlFilename2);
const data2 = new mujoco.MjData(model2); const data2 = new mujoco.MjData(model2);
try { try {
mujoco.mj_forward(model!, data!); mujoco.mj_forward(model!, data!);
@@ -1635,7 +1635,7 @@ describe('MuJoCo WASM Bindings', () => {
`; `;
writeXMLFile(tempXmlFilename, TEST_XML_TEXTURE); writeXMLFile(tempXmlFilename, TEST_XML_TEXTURE);
const model = mujoco.MjModel!.loadFromXML(tempXmlFilename); const model = mujoco.MjModel!.mj_loadXML(tempXmlFilename);
try { try {
expect(model).toBeDefined(); expect(model).toBeDefined();
expect(model!.tex_height).toEqual(new Int32Array([512])); expect(model!.tex_height).toEqual(new Int32Array([512]));
@@ -1754,7 +1754,7 @@ describe('MuJoCo WASM Bindings', () => {
</mujoco>`; </mujoco>`;
writeXMLFile(tempXmlFilename, xml); writeXMLFile(tempXmlFilename, xml);
const model = mujoco.MjModel!.loadFromXML(tempXmlFilename); const model = mujoco.MjModel!.mj_loadXML(tempXmlFilename);
try { try {
mujoco.mj_saveLastXML(tempXmlFilename, model!); mujoco.mj_saveLastXML(tempXmlFilename, model!);
const savedXmlContent = const savedXmlContent =
@@ -1781,7 +1781,7 @@ describe('MuJoCo WASM Bindings', () => {
</actuator> </actuator>
</mujoco>`; </mujoco>`;
writeXMLFile(tempXmlFilename, actuatorXml); writeXMLFile(tempXmlFilename, actuatorXml);
const model = mujoco.MjModel.loadFromXML(tempXmlFilename); const model = mujoco.MjModel.mj_loadXML(tempXmlFilename);
assertExists(model); assertExists(model);
const data = new mujoco.MjData(model); const data = new mujoco.MjData(model);
assertExists(data); assertExists(data);
+1 -1
View File
@@ -43,7 +43,7 @@ async function main() {
try { try {
console.log('Hello world!: Loading model'); console.log('Hello world!: Loading model');
model = mujoco.MjModel.loadFromXML('/working/hello.xml'); model = mujoco.MjModel.mj_loadXML('/working/hello.xml');
if (!model) { if (!model) {
throw new Error('Failed to load model'); throw new Error('Failed to load model');
} }