From d5065f867b869bfdbdb945ab9fed6552f29107e3 Mon Sep 17 00:00:00 2001 From: Alessio Quaglino Date: Mon, 17 Feb 2025 08:45:20 -0800 Subject: [PATCH] Check type correctness in all bind paths. PiperOrigin-RevId: 727877128 Change-Id: I4f4f74a1fb96d3fc5abdea531b78eb54dc53d186 --- mjx/mujoco/mjx/_src/support.py | 12 +++++++++--- mjx/mujoco/mjx/_src/support_test.py | 16 +++++++++++++--- 2 files changed, 22 insertions(+), 6 deletions(-) diff --git a/mjx/mujoco/mjx/_src/support.py b/mjx/mujoco/mjx/_src/support.py index ea777a89..bcb83c64 100644 --- a/mjx/mujoco/mjx/_src/support.py +++ b/mjx/mujoco/mjx/_src/support.py @@ -433,6 +433,12 @@ class BindData(object): self.id = ids def __getname(self, name: str): + """Get the name of the attribute and check if the type is correct.""" + if name == 'sensordata': + if self.prefix == 'sensor_': + return name + else: + raise AttributeError('sensordata is not available for this type') if name == 'ctrl': if self.prefix == 'actuator_': return name @@ -464,11 +470,11 @@ class BindData(object): idx = [] for a, n in zip(adr, num): idx.extend(a + j for j in range(n)) - return getattr(self.data, name)[idx, ...] + return getattr(self.data, self.__getname(name))[idx, ...] elif num > 1: - return getattr(self.data, name)[adr : adr + num, ...] + return getattr(self.data, self.__getname(name))[adr : adr + num, ...] else: - return getattr(self.data, name)[adr, ...] + return getattr(self.data, self.__getname(name))[adr, ...] return getattr(self.data, self.__getname(name))[self.id, ...] def set(self, name: str, value: jax.Array) -> Data: diff --git a/mjx/mujoco/mjx/_src/support_test.py b/mjx/mujoco/mjx/_src/support_test.py index e4ae3d50..34c5ba50 100644 --- a/mjx/mujoco/mjx/_src/support_test.py +++ b/mjx/mujoco/mjx/_src/support_test.py @@ -282,12 +282,22 @@ class SupportTest(parameterized.TestCase): ) # test invalid name - with self.assertRaises(AttributeError): + with self.assertRaises( + AttributeError, msg='ctrl is not available for this type' + ): print(dx.bind(mx, s.geoms).ctrl) - with self.assertRaises(AttributeError): + with self.assertRaises( + AttributeError, msg='ctrl is not available for this type' + ): print(dx.bind(mx, s.actuators).actuator_ctrl) - with self.assertRaises(AttributeError): + with self.assertRaises( + AttributeError, msg='ctrl is not available for this type' + ): print(dx.bind(mx, s.actuators).set('actuator_ctrl', [1, 2, 3])) + with self.assertRaises( + AttributeError, msg='qpos and qvel are not available for this type' + ): + print(dx.bind(mx, s.geoms).qpos) with self.assertRaises(KeyError, msg='invalid name: invalid_actuator_name'): s.actuators[0].name = 'invalid_actuator_name' print(dx.bind(mx, s.actuators).set('ctrl', [1, 2, 3]))