Migrate types in mjs structs.

PiperOrigin-RevId: 942546629
Change-Id: I46f2b24b53a43361fdc67dc32f2af1c0ce8b1de8
This commit is contained in:
Yuval Tassa
2026-07-04 11:27:00 -07:00
committed by Copybara-Service
parent 315bcfbf3a
commit 11f1da0c44
15 changed files with 339 additions and 311 deletions
@@ -304,6 +304,8 @@ def _resolve_type_name(raw_name):
return C_TO_CS_TYPE[name]
if name in _STRUCT_NAME_OVERRIDES:
return _STRUCT_NAME_OVERRIDES[name]
if name in introspect_enums.ENUMS:
return name
if name.endswith('_'):
return name
return name + '_'
@@ -19,6 +19,7 @@ from collections.abc import Sequence
from absl import app
from introspect import ast_nodes
from introspect import enums
from introspect import structs
@@ -96,6 +97,7 @@ def _value_binding_code(
fulltype = fulltype.replace('mjVisual', 'raw::MjVisual')
fulltype = fulltype.replace('mjStatistic', 'raw::MjStatistic')
element = ''
is_enum = field.name in enums.ENUMS
if field.name == 'mjsPlugin':
setter = f"""[]({rawclassname}& self, {fulltype} {varname}) {{
@@ -104,6 +106,10 @@ def _value_binding_code(
self.{fullvarname}.active = {varname}.active;
if (self.{fullvarname}.info && {varname}.info) *self.{fullvarname}.info = *{varname}.info;
}}"""
elif is_enum:
setter = f"""[]({rawclassname}& self, int {varname}) {{
self.{fullvarname}{element} = static_cast<{field.name}>({varname}){element};
}}"""
else:
setter = f"""[]({rawclassname}& self, {fulltype} {varname}) {{
self.{fullvarname}{element} = {varname}{element};
@@ -111,13 +117,13 @@ def _value_binding_code(
def_property_args = (
f'"{varname}"',
f"""[]({rawclassname}& self) -> {fulltype} {{
f"""[]({rawclassname}& self) -> {field.name if is_enum else fulltype} {{
return self.{fullvarname};
}}""",
setter,
)
if field.name not in SCALAR_TYPES:
if field.name not in SCALAR_TYPES and not is_enum:
def_property_args += ('py::return_value_policy::reference_internal',)
return f'{classname}.def_property({",".join(def_property_args)});'
+51 -51
View File
@@ -6890,7 +6890,7 @@ STRUCTS: Mapping[str, StructDecl] = dict([
fields=(
StructFieldDecl(
name='autolimits',
type=ValueType(name='mjtByte'),
type=ValueType(name='mjtBool'),
doc='infer "limited" attribute based on range',
),
StructFieldDecl(
@@ -6910,17 +6910,17 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
StructFieldDecl(
name='balanceinertia',
type=ValueType(name='mjtByte'),
type=ValueType(name='mjtBool'),
doc='automatically impose A + B >= C rule',
),
StructFieldDecl(
name='fitaabb',
type=ValueType(name='mjtByte'),
type=ValueType(name='mjtBool'),
doc='meshfit to aabb instead of inertia box',
),
StructFieldDecl(
name='degree',
type=ValueType(name='mjtByte'),
type=ValueType(name='mjtBool'),
doc='angles in radians or degrees',
),
StructFieldDecl(
@@ -6933,23 +6933,23 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
StructFieldDecl(
name='discardvisual',
type=ValueType(name='mjtByte'),
type=ValueType(name='mjtBool'),
doc='discard visual geoms in parser',
),
StructFieldDecl(
name='usethread',
type=ValueType(name='mjtByte'),
type=ValueType(name='mjtBool'),
doc='use multiple threads to speed up compiler',
),
StructFieldDecl(
name='fusestatic',
type=ValueType(name='mjtByte'),
type=ValueType(name='mjtBool'),
doc='fuse static bodies with parent',
),
StructFieldDecl(
name='inertiafromgeom',
type=ValueType(name='int'),
doc='use geom inertias (mjtInertiaFromGeom)',
type=ValueType(name='mjtInertiaFromGeom'),
doc='use geom inertias',
),
StructFieldDecl(
name='inertiagrouprange',
@@ -6961,18 +6961,18 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
StructFieldDecl(
name='saveinertial',
type=ValueType(name='mjtByte'),
type=ValueType(name='mjtBool'),
doc='save explicit inertial clause for all bodies to XML',
),
StructFieldDecl(
name='alignfree',
type=ValueType(name='int'),
type=ValueType(name='mjtBool'),
doc='align free joints with inertial frame',
),
StructFieldDecl(
name='conflict',
type=ValueType(name='int'),
doc='conflict resolution for attach (mjtConflict)',
type=ValueType(name='mjtConflict'),
doc='conflict resolution for attach',
),
StructFieldDecl(
name='LRopt',
@@ -7083,7 +7083,7 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
StructFieldDecl(
name='strippath',
type=ValueType(name='mjtByte'),
type=ValueType(name='mjtBool'),
doc='automatically strip paths from mesh files',
),
StructFieldDecl(
@@ -7192,7 +7192,7 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
StructFieldDecl(
name='hasImplicitPluginElem',
type=ValueType(name='mjtByte'),
type=ValueType(name='mjtBool'),
doc='already encountered an implicit plugin sensor/actuator',
),
StructFieldDecl(
@@ -7274,7 +7274,7 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
StructFieldDecl(
name='active',
type=ValueType(name='mjtByte'),
type=ValueType(name='mjtBool'),
doc='is the plugin active',
),
StructFieldDecl(
@@ -7370,7 +7370,7 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
StructFieldDecl(
name='mocap',
type=ValueType(name='mjtByte'),
type=ValueType(name='mjtBool'),
doc='is this a mocap body',
),
StructFieldDecl(
@@ -7397,7 +7397,7 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
StructFieldDecl(
name='explicitinertial',
type=ValueType(name='mjtByte'),
type=ValueType(name='mjtBool'),
doc='whether to save the body with explicit inertial clause',
),
StructFieldDecl(
@@ -7503,8 +7503,8 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
StructFieldDecl(
name='align',
type=ValueType(name='int'),
doc='align free joint with body com (mjtAlignFree)',
type=ValueType(name='mjtAlignFree'),
doc='align free joint with body com',
),
StructFieldDecl(
name='stiffness',
@@ -7529,8 +7529,8 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
StructFieldDecl(
name='limited',
type=ValueType(name='int'),
doc='does joint have limits (mjtLimited)',
type=ValueType(name='mjtLimited'),
doc='does joint have limits',
),
StructFieldDecl(
name='range',
@@ -7563,8 +7563,8 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
StructFieldDecl(
name='actfrclimited',
type=ValueType(name='int'),
doc='are actuator forces on joint limited (mjtLimited)',
type=ValueType(name='mjtLimited'),
doc='are actuator forces on joint limited',
),
StructFieldDecl(
name='actfrcrange',
@@ -7615,7 +7615,7 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
StructFieldDecl(
name='actgravcomp',
type=ValueType(name='mjtByte'),
type=ValueType(name='mjtBool'),
doc='is gravcomp force applied via actuators',
),
StructFieldDecl(
@@ -8104,7 +8104,7 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
StructFieldDecl(
name='active',
type=ValueType(name='mjtByte'),
type=ValueType(name='mjtBool'),
doc='is light active',
),
StructFieldDecl(
@@ -8121,7 +8121,7 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
StructFieldDecl(
name='castshadow',
type=ValueType(name='mjtByte'),
type=ValueType(name='mjtBool'),
doc='does light cast shadows',
),
StructFieldDecl(
@@ -8186,7 +8186,7 @@ STRUCTS: Mapping[str, StructDecl] = dict([
type=PointerType(
inner_type=ValueType(name='mjString'),
),
doc='message appended to compiler errorsx',
doc='message appended to compiler errors',
),
),
)),
@@ -8281,17 +8281,17 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
StructFieldDecl(
name='internal',
type=ValueType(name='mjtByte'),
type=ValueType(name='mjtBool'),
doc='enable internal collisions',
),
StructFieldDecl(
name='flatskin',
type=ValueType(name='mjtByte'),
type=ValueType(name='mjtBool'),
doc='render flex skin with flat shading',
),
StructFieldDecl(
name='selfcollide',
type=ValueType(name='int'),
type=ValueType(name='mjtFlexSelf'),
doc='mode for flex self collision',
),
StructFieldDecl(
@@ -8487,12 +8487,12 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
StructFieldDecl(
name='smoothnormal',
type=ValueType(name='mjtByte'),
type=ValueType(name='mjtBool'),
doc='do not exclude large-angle faces from normals',
),
StructFieldDecl(
name='needsdf',
type=ValueType(name='mjtByte'),
type=ValueType(name='mjtBool'),
doc='compute sdf from mesh',
),
StructFieldDecl(
@@ -8761,13 +8761,13 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
StructFieldDecl(
name='builtin',
type=ValueType(name='int'),
doc='builtin type (mjtBuiltin)',
type=ValueType(name='mjtBuiltin'),
doc='builtin type',
),
StructFieldDecl(
name='mark',
type=ValueType(name='int'),
doc='mark type (mjtMark)',
type=ValueType(name='mjtMark'),
doc='mark type',
),
StructFieldDecl(
name='rgb1',
@@ -8859,12 +8859,12 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
StructFieldDecl(
name='hflip',
type=ValueType(name='mjtByte'),
type=ValueType(name='mjtBool'),
doc='horizontal flip',
),
StructFieldDecl(
name='vflip',
type=ValueType(name='mjtByte'),
type=ValueType(name='mjtBool'),
doc='vertical flip',
),
StructFieldDecl(
@@ -8897,7 +8897,7 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
StructFieldDecl(
name='texuniform',
type=ValueType(name='mjtByte'),
type=ValueType(name='mjtBool'),
doc='make texture cube uniform',
),
StructFieldDecl(
@@ -9099,7 +9099,7 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
StructFieldDecl(
name='active',
type=ValueType(name='mjtByte'),
type=ValueType(name='mjtBool'),
doc='is equality initially active',
),
StructFieldDecl(
@@ -9210,12 +9210,12 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
StructFieldDecl(
name='limited',
type=ValueType(name='int'),
doc='does tendon have limits (mjtLimited)',
type=ValueType(name='mjtLimited'),
doc='does tendon have limits',
),
StructFieldDecl(
name='actfrclimited',
type=ValueType(name='int'),
type=ValueType(name='mjtLimited'),
doc='does tendon have actuator force limits',
),
StructFieldDecl(
@@ -9380,7 +9380,7 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
StructFieldDecl(
name='actearly',
type=ValueType(name='mjtByte'),
type=ValueType(name='mjtBool'),
doc='apply next activations to qfrc',
),
StructFieldDecl(
@@ -9450,8 +9450,8 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
StructFieldDecl(
name='ctrllimited',
type=ValueType(name='int'),
doc='are control limits defined (mjtLimited)',
type=ValueType(name='mjtLimited'),
doc='are control limits defined',
),
StructFieldDecl(
name='ctrlrange',
@@ -9463,8 +9463,8 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
StructFieldDecl(
name='forcelimited',
type=ValueType(name='int'),
doc='are force limits defined (mjtLimited)',
type=ValueType(name='mjtLimited'),
doc='are force limits defined',
),
StructFieldDecl(
name='forcerange',
@@ -9476,8 +9476,8 @@ STRUCTS: Mapping[str, StructDecl] = dict([
),
StructFieldDecl(
name='actlimited',
type=ValueType(name='int'),
doc='are activation limits defined (mjtLimited)',
type=ValueType(name='mjtLimited'),
doc='are activation limits defined',
),
StructFieldDecl(
name='actrange',
+1 -1
View File
@@ -777,7 +777,7 @@ PYBIND11_MODULE(_specs, m, pybind11::mod_gil_not_used()) {
std::string key = py::str(item.first);
if (key == "align") {
try {
out->align = kwargs["align"].cast<int>();
out->align = static_cast<mjtAlignFree>(kwargs["align"].cast<int>());
} catch (const py::cast_error& e) {
throw pybind11::value_error("align is the wrong type.");
}