Replace name attribute with setters and getters in the C API of mjSpec.

PiperOrigin-RevId: 778816015
Change-Id: Ieffb7a5bce37e887ff009f9a02d2434435e79ca1
This commit is contained in:
Alessio Quaglino
2025-07-03 02:29:49 -07:00
committed by Copybara-Service
parent 1c5d47c94b
commit 594e17074a
22 changed files with 290 additions and 400 deletions
@@ -341,9 +341,14 @@ def generate_add() -> None:
else:
return '', '', ''
code_field = ''
set_types = []
names = []
if key == 'mjsPlugin':
code_field = ''
set_types = []
names = []
else:
code_field = 'set_name("name", out->element);'
set_types = ['name']
names = ['name']
for field in structs.STRUCTS[key].fields:
line, set_type, name = _field(field)
if line:
@@ -571,6 +576,21 @@ def generate_add() -> None:
}
};
"""
elif t == 'name':
code += """\n
auto set_name = [&kwargs](const char* str, raw::MjsElement* el) {
if (kwargs.contains(str)) {
try {
std::string name = kwargs[str].cast<std::string>();
mjs_setName(el, name.c_str());
} catch (const py::cast_error &e) {
throw pybind11::value_error(std::string(str) + " should be a string.");
}
}
};
"""
else:
raise NotImplementedError(f'Unsupported set type: {t} in {key}')
code += code_field
code += f"""\n
@@ -644,6 +664,25 @@ def generate_id() -> None:
print(code)
def generate_name() -> None:
"""Generate name functions."""
for key, _, _, _, _ in SPECS + [('mjsDefault', '', '', '', '')]:
if key == 'mjsPlugin':
continue
elem = key.removeprefix('mjs')
titlecase = 'Mjs' + elem
code = f"""\n
{key}.def_property("name",
[](raw::{titlecase}& self) -> std::string* {{
return mjs_getName(self.element);
}},
[](raw::{titlecase}& self, std::string& name) -> void {{
mjs_setName(self.element, name.c_str());
}}, py::return_value_policy::reference_internal);
"""
print(code)
def main(argv: Sequence[str]) -> None:
if len(argv) > 1:
raise app.UsageError('Too many command-line arguments.')
@@ -652,6 +691,7 @@ def main(argv: Sequence[str]) -> None:
generate_find()
generate_signature()
generate_id()
generate_name()
if __name__ == '__main__':
+36
View File
@@ -10191,6 +10191,26 @@ FUNCTIONS: Mapping[str, FunctionDecl] = dict([
),
doc="Return spec's next element; return NULL if element is last.",
)),
('mjs_setName',
FunctionDecl(
name='mjs_setName',
return_type=ValueType(name='void'),
parameters=(
FunctionParameterDecl(
name='element',
type=PointerType(
inner_type=ValueType(name='mjsElement'),
),
),
FunctionParameterDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='char', is_const=True),
),
),
),
doc="Set element's name.",
)),
('mjs_setBuffer',
FunctionDecl(
name='mjs_setBuffer',
@@ -10439,6 +10459,22 @@ FUNCTIONS: Mapping[str, FunctionDecl] = dict([
),
doc='Set plugin attributes.',
)),
('mjs_getName',
FunctionDecl(
name='mjs_getName',
return_type=PointerType(
inner_type=ValueType(name='mjString'),
),
parameters=(
FunctionParameterDecl(
name='element',
type=PointerType(
inner_type=ValueType(name='mjsElement'),
),
),
),
doc="Get element's name.",
)),
('mjs_getString',
FunctionDecl(
name='mjs_getString',
-168
View File
@@ -8222,13 +8222,6 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
doc='element type',
),
StructFieldDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='mjString'),
),
doc='name',
),
StructFieldDecl(
name='childclass',
type=PointerType(
@@ -8347,13 +8340,6 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
doc='element type',
),
StructFieldDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='mjString'),
),
doc='name',
),
StructFieldDecl(
name='childclass',
type=PointerType(
@@ -8403,13 +8389,6 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
doc='element type',
),
StructFieldDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='mjString'),
),
doc='name',
),
StructFieldDecl(
name='type',
type=ValueType(name='mjtJoint'),
@@ -8575,13 +8554,6 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
doc='element type',
),
StructFieldDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='mjString'),
),
doc='name',
),
StructFieldDecl(
name='type',
type=ValueType(name='mjtGeom'),
@@ -8783,13 +8755,6 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
doc='element type',
),
StructFieldDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='mjString'),
),
doc='name',
),
StructFieldDecl(
name='pos',
type=ArrayType(
@@ -8880,13 +8845,6 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
doc='element type',
),
StructFieldDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='mjString'),
),
doc='name',
),
StructFieldDecl(
name='pos',
type=ArrayType(
@@ -9019,13 +8977,6 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
doc='element type',
),
StructFieldDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='mjString'),
),
doc='name',
),
StructFieldDecl(
name='pos',
type=ArrayType(
@@ -9154,13 +9105,6 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
doc='element type',
),
StructFieldDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='mjString'),
),
doc='name',
),
StructFieldDecl(
name='contype',
type=ValueType(name='int'),
@@ -9380,13 +9324,6 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
doc='element type',
),
StructFieldDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='mjString'),
),
doc='name',
),
StructFieldDecl(
name='content_type',
type=PointerType(
@@ -9501,13 +9438,6 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
doc='element type',
),
StructFieldDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='mjString'),
),
doc='name',
),
StructFieldDecl(
name='content_type',
type=PointerType(
@@ -9568,13 +9498,6 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
doc='element type',
),
StructFieldDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='mjString'),
),
doc='name',
),
StructFieldDecl(
name='file',
type=PointerType(
@@ -9684,13 +9607,6 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
doc='element type',
),
StructFieldDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='mjString'),
),
doc='name',
),
StructFieldDecl(
name='type',
type=ValueType(name='mjtTexture'),
@@ -9830,13 +9746,6 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
doc='element type',
),
StructFieldDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='mjString'),
),
doc='name',
),
StructFieldDecl(
name='textures',
type=PointerType(
@@ -9916,13 +9825,6 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
doc='element type',
),
StructFieldDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='mjString'),
),
doc='name',
),
StructFieldDecl(
name='geomname1',
type=PointerType(
@@ -10005,13 +9907,6 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
doc='element type',
),
StructFieldDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='mjString'),
),
doc='name',
),
StructFieldDecl(
name='bodyname1',
type=PointerType(
@@ -10047,13 +9942,6 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
doc='element type',
),
StructFieldDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='mjString'),
),
doc='name',
),
StructFieldDecl(
name='type',
type=ValueType(name='mjtEq'),
@@ -10128,13 +10016,6 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
doc='element type',
),
StructFieldDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='mjString'),
),
doc='name',
),
StructFieldDecl(
name='stiffness',
type=ValueType(name='double'),
@@ -10300,13 +10181,6 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
doc='element type',
),
StructFieldDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='mjString'),
),
doc='name',
),
StructFieldDecl(
name='gaintype',
type=ValueType(name='mjtGain'),
@@ -10485,13 +10359,6 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
doc='element type',
),
StructFieldDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='mjString'),
),
doc='name',
),
StructFieldDecl(
name='type',
type=ValueType(name='mjtSensor'),
@@ -10579,13 +10446,6 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
doc='element type',
),
StructFieldDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='mjString'),
),
doc='name',
),
StructFieldDecl(
name='data',
type=PointerType(
@@ -10619,13 +10479,6 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
doc='element type',
),
StructFieldDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='mjString'),
),
doc='name',
),
StructFieldDecl(
name='data',
type=PointerType(
@@ -10654,13 +10507,6 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
doc='element type',
),
StructFieldDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='mjString'),
),
doc='name',
),
StructFieldDecl(
name='objtype',
type=PointerType(
@@ -10703,13 +10549,6 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
doc='element type',
),
StructFieldDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='mjString'),
),
doc='name',
),
StructFieldDecl(
name='time',
type=ValueType(name='double'),
@@ -10778,13 +10617,6 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
doc='element type',
),
StructFieldDecl(
name='name',
type=PointerType(
inner_type=ValueType(name='mjString'),
),
doc='class name',
),
StructFieldDecl(
name='joint',
type=PointerType(
+24 -12
View File
@@ -521,7 +521,8 @@ PYBIND11_MODULE(_specs, m) {
}
} else if (key == "name") {
try {
*out->name = kwargs["name"].cast<std::string>();
mjs_setName(out->element,
kwargs["name"].cast<std::string>().c_str());
} catch (const py::cast_error& e) {
throw pybind11::value_error("name is the wrong type.");
}
@@ -552,7 +553,8 @@ PYBIND11_MODULE(_specs, m) {
},
[](raw::MjsBody& self, raw::MjsDefault& default_) -> void {
mjs_setDefault(self.element, &default_);
});
},
py::return_value_policy::reference_internal);
mjsBody.def(
"find_all",
[](raw::MjsBody& self, mjtObj objtype) -> py::list {
@@ -819,7 +821,8 @@ PYBIND11_MODULE(_specs, m) {
},
[](raw::MjsGeom& self, raw::MjsDefault& default_) -> void {
mjs_setDefault(self.element, &default_);
});
},
py::return_value_policy::reference_internal);
mjsGeom.def_property_readonly(
"frame",
[](raw::MjsGeom& self) -> raw::MjsFrame* {
@@ -849,7 +852,8 @@ PYBIND11_MODULE(_specs, m) {
},
[](raw::MjsJoint& self, raw::MjsDefault& default_) -> void {
mjs_setDefault(self.element, &default_);
});
},
py::return_value_policy::reference_internal);
mjsJoint.def_property_readonly(
"frame",
[](raw::MjsJoint& self) -> raw::MjsFrame* {
@@ -879,7 +883,8 @@ PYBIND11_MODULE(_specs, m) {
},
[](raw::MjsSite& self, raw::MjsDefault& default_) -> void {
mjs_setDefault(self.element, &default_);
});
},
py::return_value_policy::reference_internal);
mjsSite.def(
"attach_body",
[](raw::MjsSite& self, raw::MjsBody& body,
@@ -926,7 +931,8 @@ PYBIND11_MODULE(_specs, m) {
},
[](raw::MjsCamera& self, raw::MjsDefault& default_) -> void {
mjs_setDefault(self.element, &default_);
});
},
py::return_value_policy::reference_internal);
mjsCamera.def_property_readonly(
"frame",
[](raw::MjsCamera& self) -> raw::MjsFrame* {
@@ -956,7 +962,8 @@ PYBIND11_MODULE(_specs, m) {
},
[](raw::MjsLight& self, raw::MjsDefault& default_) -> void {
mjs_setDefault(self.element, &default_);
});
},
py::return_value_policy::reference_internal);
mjsLight.def_property_readonly(
"frame",
[](raw::MjsLight& self) -> raw::MjsFrame* {
@@ -975,7 +982,8 @@ PYBIND11_MODULE(_specs, m) {
},
[](raw::MjsMaterial& self, raw::MjsDefault& default_) -> void {
mjs_setDefault(self.element, &default_);
});
},
py::return_value_policy::reference_internal);
// ============================= MJSMESH =====================================
mjSpec.def("delete", [](MjSpec& self, raw::MjsMesh& obj) {
@@ -988,7 +996,8 @@ PYBIND11_MODULE(_specs, m) {
},
[](raw::MjsMesh& self, raw::MjsDefault& default_) -> void {
mjs_setDefault(self.element, &default_);
});
},
py::return_value_policy::reference_internal);
// ============================= MJSPAIR =====================================
mjSpec.def("delete", [](MjSpec& self, raw::MjsPair& obj) {
@@ -1001,7 +1010,8 @@ PYBIND11_MODULE(_specs, m) {
},
[](raw::MjsPair& self, raw::MjsDefault& default_) -> void {
mjs_setDefault(self.element, &default_);
});
},
py::return_value_policy::reference_internal);
// ============================= MJSEQUAL ====================================
mjSpec.def("delete", [](MjSpec& self, raw::MjsEquality& obj) {
@@ -1014,7 +1024,8 @@ PYBIND11_MODULE(_specs, m) {
},
[](raw::MjsEquality& self, raw::MjsDefault& default_) -> void {
mjs_setDefault(self.element, &default_);
});
},
py::return_value_policy::reference_internal);
// ============================= MJSACTUATOR =================================
mjSpec.def("delete", [](MjSpec& self, raw::MjsActuator& obj) {
@@ -1027,7 +1038,8 @@ PYBIND11_MODULE(_specs, m) {
},
[](raw::MjsActuator& self, raw::MjsDefault& default_) -> void {
mjs_setDefault(self.element, &default_);
});
},
py::return_value_policy::reference_internal);
mjsActuator.def("set_to_motor", [](raw::MjsActuator* self) {
std::string err = mjs_setToMotor(self);
if (!err.empty()) {