Refactor MuJoCo WASM struct codegen.

PiperOrigin-RevId: 837530063
Change-Id: I57518b2fbb536f356a5e15839d8e17f455bccde9
This commit is contained in:
Matias Manevi
2025-11-27 07:41:42 -08:00
committed by Copybara-Service
parent 5163dfa823
commit 9ca1598b23
4 changed files with 257 additions and 467 deletions
+17 -202
View File
@@ -203,58 +203,6 @@ class FunctionUtilsTest(absltest.TestCase):
class StructConstructorCodeBuilderTest(absltest.TestCase):
def test_constructor_code_with_default_function(self):
wrapped_structs = structs.generate_wasm_bindings(["mjLROpt"])
self.assertEqual(
wrapped_structs["mjLROpt"].wrapped_source,
"""
MjLROpt::MjLROpt(mjLROpt *ptr) : ptr_(ptr) {}
MjLROpt::~MjLROpt() {
if (owned_ && ptr_) {
delete ptr_;
}
}
MjLROpt::MjLROpt() : ptr_(new mjLROpt) {
owned_ = true;
mj_defaultLROpt(ptr_);
}
MjLROpt::MjLROpt(const MjLROpt &other) : MjLROpt() {
*ptr_ = *other.get();
}
MjLROpt& MjLROpt::operator=(const MjLROpt &other) {
if (this == &other) {
return *this;
}
*ptr_ = *other.get();
return *this;
}
std::unique_ptr<MjLROpt> MjLROpt::copy() {
return std::make_unique<MjLROpt>(*this);
}
mjLROpt* MjLROpt::get() const {
return ptr_;
}
void MjLROpt::set(mjLROpt* ptr) {
ptr_ = ptr;
}
""".strip(),
)
def test_constructor_code_without_default_function(self):
wrapped_structs = structs.generate_wasm_bindings(["mjsElement"])
self.assertEqual(
wrapped_structs["mjsElement"].wrapped_source,
"""
MjsElement::MjsElement(mjsElement *ptr) : ptr_(ptr) {}
mjsElement* MjsElement::get() const {
return ptr_;
}
void MjsElement::set(mjsElement* ptr) {
ptr_ = ptr;
}
""".strip(),
)
def test_constructor_code_with_fields_with_init(self):
field_with_init = ast_nodes.StructFieldDecl(
name="element",
@@ -354,7 +302,7 @@ class StructFieldCodeBuilderTest(absltest.TestCase):
doc="number of geoms",
)
self.assertEqual(
structs.build_primitive_type_definition(field),
structs._generate_field_data(field, "ngeom").declaration,
"""
int ngeom() const {
return ptr_->ngeom;
@@ -375,9 +323,7 @@ void set_ngeom(int value) {
array_extent=("ngeom", 4),
)
self.assertEqual(
structs.build_memory_view_definition(
field, "ptr_->ngeom * 4", "ptr_->geom_rgba"
),
structs._generate_field_data(field, "geom_rgba").declaration,
"""
emscripten::val geom_rgba() const {
return emscripten::val(emscripten::typed_memory_view(ptr_->ngeom * 4, ptr_->geom_rgba));
@@ -394,7 +340,7 @@ emscripten::val geom_rgba() const {
doc="rgba when material is omitted",
)
self.assertEqual(
structs.build_string_field_definition(field),
structs._generate_field_data(field, "MjString").declaration,
"""
mjString string_field() const {
return (ptr_ && ptr_->string_field) ? *(ptr_->string_field) : "";
@@ -416,7 +362,7 @@ void set_string_field(const mjString& value) {
doc="",
)
self.assertEqual(
structs.build_mjvec_pointer_definition(field, "mjDoubleVec"),
structs._generate_field_data(field, "mjDoubleVec").declaration,
"""
mjDoubleVec &vector_field() const {
return *(ptr_->vector_field);
@@ -432,174 +378,43 @@ mjDoubleVec &vector_field() const {
doc="",
)
self.assertEqual(
structs.build_mjvec_pointer_definition(field, "mjByteVec"),
structs._generate_field_data(field, "mjByteVec").declaration,
"""
std::vector<uint8_t> &vector_field() const {
return *(reinterpret_cast<std::vector<uint8_t>*>(ptr_->vector_field));
}""".strip(),
)
def test_simple_property_binding(self):
def test_get_property_binding(self):
field = ast_nodes.StructFieldDecl(
name="ngeom",
type=ast_nodes.ValueType(name="int"),
doc="number of geoms",
)
self.assertEqual(
structs.build_simple_property_binding(field, "MjModel"),
structs._get_property_binding(field, "MjModel"),
'.property("ngeom", &MjModel::ngeom)',
)
def test_simple_property_binding_with_setter(self):
def test_get_property_binding_with_setter(self):
field = ast_nodes.StructFieldDecl(
name="ngeom",
type=ast_nodes.ValueType(name="int"),
doc="",
)
self.assertEqual(
structs.build_simple_property_binding(field, "MjModel", True),
structs._get_property_binding(field, "MjModel", True),
'.property("ngeom", &MjModel::ngeom, &MjModel::set_ngeom)',
)
def test_simple_property_binding_with_return_value_policy_as_ref(self):
def test_get_property_binding_with_return_value_policy_as_ref(self):
field = ast_nodes.StructFieldDecl(
name="ngeom",
type=ast_nodes.ValueType(name="int"),
doc="",
)
self.assertEqual(
structs.build_simple_property_binding(
field,
"MjModel",
add_setter=True,
add_return_value_policy_as_ref=True,
),
'.property("ngeom", &MjModel::ngeom, &MjModel::set_ngeom, reference())',
)
class StructFieldCodeBuilderTest(absltest.TestCase):
def test_primitive_type_definition(self):
field = ast_nodes.StructFieldDecl(
name="ngeom",
type=ast_nodes.ValueType(name="int"),
doc="number of geoms",
)
self.assertEqual(
structs._generate_field_data(field, "ngeom").definition,
"""
int ngeom() const {
return ptr_->ngeom;
}
void set_ngeom(int value) {
ptr_->ngeom = value;
}
""".strip(),
)
def test_memory_view_definition(self):
field = ast_nodes.StructFieldDecl(
name="geom_rgba",
type=ast_nodes.PointerType(
inner_type=ast_nodes.ValueType(name="float"),
),
doc="rgba when material is omitted",
array_extent=("ngeom", 4),
)
self.assertEqual(
structs._generate_field_data(field, "geom_rgba").definition,
"""
emscripten::val geom_rgba() const {
return emscripten::val(emscripten::typed_memory_view(ptr_->ngeom * 4, ptr_->geom_rgba));
}
""".strip(),
)
def test_string_field_definition(self):
field = ast_nodes.StructFieldDecl(
name="string_field",
type=ast_nodes.PointerType(
inner_type=ast_nodes.ValueType(name="mjString"),
),
doc="rgba when material is omitted",
)
self.assertEqual(
structs._generate_field_data(field, "MjString").definition,
"""
mjString string_field() const {
return (ptr_ && ptr_->string_field) ? *(ptr_->string_field) : "";
}
void set_string_field(const mjString& value) {
if (ptr_ && ptr_->string_field) {
*(ptr_->string_field) = value;
}
}
""".strip(),
)
def test_mjvec_pointer_definition(self):
field = ast_nodes.StructFieldDecl(
name="vector_field",
type=ast_nodes.PointerType(
inner_type=ast_nodes.ValueType(name="mjDoubleVec"),
),
doc="",
)
self.assertEqual(
structs._generate_field_data(field, "mjDoubleVec").definition,
"""
mjDoubleVec &vector_field() const {
return *(ptr_->vector_field);
}""".strip(),
)
def test_mjbyte_vec_pointer_definition(self):
field = ast_nodes.StructFieldDecl(
name="vector_field",
type=ast_nodes.PointerType(
inner_type=ast_nodes.ValueType(name="mjByteVec"),
),
doc="",
)
self.assertEqual(
structs._generate_field_data(field, "mjByteVec").definition,
"""
std::vector<uint8_t> &vector_field() const {
return *(reinterpret_cast<std::vector<uint8_t>*>(ptr_->vector_field));
}""".strip(),
)
def test_simple_property_binding(self):
field = ast_nodes.StructFieldDecl(
name="ngeom",
type=ast_nodes.ValueType(name="int"),
doc="number of geoms",
)
self.assertEqual(
structs._simple_property_binding(field, "MjModel"),
'.property("ngeom", &MjModel::ngeom)',
)
def test_simple_property_binding_with_setter(self):
field = ast_nodes.StructFieldDecl(
name="ngeom",
type=ast_nodes.ValueType(name="int"),
doc="",
)
self.assertEqual(
structs._simple_property_binding(field, "MjModel", True),
'.property("ngeom", &MjModel::ngeom, &MjModel::set_ngeom)',
)
def test_simple_property_binding_with_return_value_policy_as_ref(self):
field = ast_nodes.StructFieldDecl(
name="ngeom",
type=ast_nodes.ValueType(name="int"),
doc="",
)
self.assertEqual(
structs._simple_property_binding(
structs._get_property_binding(
field,
"MjModel",
setter=True,
@@ -621,7 +436,7 @@ class StructFieldHandlerTest(absltest.TestCase):
wrapped_field_data = structs._generate_field_data(field_scalar, "MjModel")
self.assertEqual(
wrapped_field_data.definition,
wrapped_field_data.declaration,
"""
int ngeom() const {
return ptr_->ngeom;
@@ -649,7 +464,7 @@ void set_ngeom(int value) {
wrapped_field_data = structs._generate_field_data(field, "MjModel")
self.assertEqual(
wrapped_field_data.definition,
wrapped_field_data.declaration,
("""
emscripten::val geom_rgba() const {
return emscripten::val(emscripten::typed_memory_view(ptr_->ngeom * 4, ptr_->geom_rgba));
@@ -674,7 +489,7 @@ emscripten::val geom_rgba() const {
wrapped_field_data = structs._generate_field_data(field, "MjData")
self.assertEqual(
wrapped_field_data.definition,
wrapped_field_data.declaration,
("""
emscripten::val buffer() const {
return emscripten::val(emscripten::typed_memory_view(model->nbuffer, static_cast<uint8_t*>(ptr_->buffer)));
@@ -693,7 +508,7 @@ emscripten::val buffer() const {
doc="",
)
wrapped_field_data = structs._generate_field_data(field, "MjsTexture")
self.assertEqual(wrapped_field_data.definition, "MjsElement element;")
self.assertEqual(wrapped_field_data.declaration, "MjsElement element;")
self.assertEqual(
wrapped_field_data.binding,
@@ -717,7 +532,7 @@ emscripten::val buffer() const {
wrapped_field_data = structs._generate_field_data(field, "MjOption")
self.assertEqual(
wrapped_field_data.definition,
wrapped_field_data.declaration,
("""
emscripten::val gravity() const {
return emscripten::val(emscripten::typed_memory_view(3, ptr_->gravity));
@@ -741,7 +556,7 @@ emscripten::val gravity() const {
)
wrapped_field_data = structs._generate_field_data(field, "MjModel")
self.assertEqual(
wrapped_field_data.definition,
wrapped_field_data.declaration,
"""
emscripten::val multi_dim_array() const {
return emscripten::val(emscripten::typed_memory_view(12, reinterpret_cast<float*>(ptr_->multi_dim_array)));