Refactor MuJoCo WASM struct codegen.
PiperOrigin-RevId: 837530063 Change-Id: I57518b2fbb536f356a5e15839d8e17f455bccde9
This commit is contained in:
committed by
Copybara-Service
parent
5163dfa823
commit
9ca1598b23
@@ -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)));
|
||||
|
||||
Reference in New Issue
Block a user