Python bindings: Give a detailed error message when passing non-existent name.
List all valid names in the error message. Fix an off-by-one error when accessing elements by numeric index (rather than name). PiperOrigin-RevId: 457702283 Change-Id: I6e4f1518b4ecd9ca91ac9b0bf6b0963fc0a9bd18
This commit is contained in:
committed by
Copybara-Service
parent
a077db1b91
commit
e14cf7ddf1
@@ -163,6 +163,46 @@ class MuJoCoBindingsTest(parameterized.TestCase):
|
||||
self.data.actuator('myactuator').ctrl = 7
|
||||
np.testing.assert_array_equal(self.data.ctrl[actuator_id], [7])
|
||||
|
||||
def test_named_indexing_invalid_names_in_model(self):
|
||||
with self.assertRaisesRegex(
|
||||
KeyError,
|
||||
r"Invalid name 'badgeom'\. Valid names: \['mybox', 'myplane'\]"):
|
||||
self.model.geom('badgeom')
|
||||
|
||||
def test_named_indexing_no_name_argument_in_model(self):
|
||||
with self.assertRaisesRegex(
|
||||
KeyError,
|
||||
r"Invalid name ''\. Valid names: \['myball', 'myfree', 'myhinge'\]"):
|
||||
self.model.joint()
|
||||
|
||||
def test_named_indexing_invalid_names_in_data(self):
|
||||
with self.assertRaisesRegex(
|
||||
KeyError,
|
||||
r"Invalid name 'badgeom'\. Valid names: \['mybox', 'myplane'\]"):
|
||||
self.data.geom('badgeom')
|
||||
|
||||
def test_named_indexing_no_name_argument_in_model(self):
|
||||
with self.assertRaisesRegex(
|
||||
KeyError,
|
||||
r"Invalid name ''\. Valid names: \['myball', 'myfree', 'myhinge'\]"):
|
||||
self.data.jnt()
|
||||
|
||||
def test_named_indexing_invalid_index_in_model(self):
|
||||
with self.assertRaisesRegex(
|
||||
IndexError, r'Invalid index 3\. Valid indices from 0 to 2'):
|
||||
self.model.geom(3)
|
||||
with self.assertRaisesRegex(
|
||||
IndexError, r'Invalid index -1\. Valid indices from 0 to 2'):
|
||||
self.model.geom(-1)
|
||||
|
||||
def test_named_indexing_invalid_index_in_data(self):
|
||||
with self.assertRaisesRegex(
|
||||
IndexError, r'Invalid index 3\. Valid indices from 0 to 2'):
|
||||
self.data.geom(3)
|
||||
with self.assertRaisesRegex(
|
||||
IndexError, r'Invalid index -1\. Valid indices from 0 to 2'):
|
||||
self.data.geom(-1)
|
||||
|
||||
def test_named_indexing_geom_size(self):
|
||||
box_id = mujoco.mj_name2id(self.model, mujoco.mjtObj.mjOBJ_GEOM, 'mybox')
|
||||
self.assertIs(self.model.geom('mybox'), self.model.geom(box_id))
|
||||
|
||||
+54
-24
@@ -12,6 +12,7 @@
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
#include <algorithm>
|
||||
#include <memory>
|
||||
#include <optional>
|
||||
#include <string>
|
||||
@@ -154,8 +155,8 @@ MjModelIndexer::MjModelIndexer(raw::MjModel* m, py::handle owner)
|
||||
|
||||
#define XGROUP(MjModelFieldGroupedViews, field, nfield, FIELD_XMACROS) \
|
||||
MjModelFieldGroupedViews& MjModelIndexer::field(int i) { \
|
||||
if (i > field##_.size()) { \
|
||||
throw py::index_error("index out of range"); \
|
||||
if (i >= field##_.size() || i < 0) { \
|
||||
throw py::index_error(IndexErrorMessage(i, field##_.size())); \
|
||||
} \
|
||||
auto& indexer = field##_[i]; \
|
||||
if (!indexer.has_value()) { \
|
||||
@@ -169,11 +170,11 @@ MJMODEL_VIEW_GROUPS
|
||||
#define XGROUP(MjModelFieldGroupedViews, field, nfield, FIELD_XMACROS) \
|
||||
MjModelFieldGroupedViews& MjModelIndexer::field##_by_name( \
|
||||
std::string_view name) { \
|
||||
try { \
|
||||
return field(name_to_id_.field.at(name)); \
|
||||
} catch (...) { \
|
||||
throw py::key_error(std::string(name)); \
|
||||
auto item = name_to_id_.field.find(name); \
|
||||
if (item == name_to_id_.field.end()) { \
|
||||
throw py::key_error(KeyErrorMessage(name_to_id_.field, name)); \
|
||||
} \
|
||||
return field(item->second); \
|
||||
}
|
||||
MJMODEL_VIEW_GROUPS
|
||||
#undef XGROUP
|
||||
@@ -190,28 +191,27 @@ MjDataIndexer::MjDataIndexer(raw::MjData* d, const MjDataMetadata* m,
|
||||
#undef XGROUP
|
||||
{}
|
||||
|
||||
#define XGROUP(MjDataGroupedViews, field, nfield, FIELD_XMACROS) \
|
||||
MjDataGroupedViews& MjDataIndexer::field(int i) { \
|
||||
if (i > field##_.size()) { \
|
||||
throw py::index_error("index out of range"); \
|
||||
} \
|
||||
auto& indexer = field##_[i]; \
|
||||
if (!indexer.has_value()) { \
|
||||
indexer.emplace(i, d_, m_, owner_); \
|
||||
} \
|
||||
return *indexer; \
|
||||
#define XGROUP(MjDataGroupedViews, field, nfield, FIELD_XMACROS) \
|
||||
MjDataGroupedViews& MjDataIndexer::field(int i) { \
|
||||
if (i >= field##_.size() || i < 0) { \
|
||||
throw py::index_error(IndexErrorMessage(i, field##_.size())); \
|
||||
} \
|
||||
auto& indexer = field##_[i]; \
|
||||
if (!indexer.has_value()) { \
|
||||
indexer.emplace(i, d_, m_, owner_); \
|
||||
} \
|
||||
return *indexer; \
|
||||
}
|
||||
MJDATA_VIEW_GROUPS
|
||||
#undef XGROUP
|
||||
|
||||
#define XGROUP(MjDataGroupedViews, field, nfield, FIELD_XMACROS) \
|
||||
MjDataGroupedViews& MjDataIndexer::field##_by_name( \
|
||||
std::string_view name) { \
|
||||
try { \
|
||||
return field(name_to_id_.field.at(name)); \
|
||||
} catch (...) { \
|
||||
throw py::key_error(std::string(name)); \
|
||||
} \
|
||||
#define XGROUP(MjDataGroupedViews, field, nfield, FIELD_XMACROS) \
|
||||
MjDataGroupedViews& MjDataIndexer::field##_by_name(std::string_view name) { \
|
||||
auto item = name_to_id_.field.find(name); \
|
||||
if (item == name_to_id_.field.end()) { \
|
||||
throw py::key_error(KeyErrorMessage(name_to_id_.field, name)); \
|
||||
} \
|
||||
return field(item->second); \
|
||||
}
|
||||
MJDATA_VIEW_GROUPS
|
||||
#undef XGROUP
|
||||
@@ -369,4 +369,34 @@ MJDATA_TENDON
|
||||
#undef MJ_M
|
||||
#define MJ_M(n) n
|
||||
|
||||
// Returns an error message when a non-existent name is requested from an
|
||||
// indexer, which includes all valid names.
|
||||
std::string KeyErrorMessage(const NameToIDMap& map, std::string_view name) {
|
||||
// Make a sorted list of valid names
|
||||
std::vector<std::string_view> valid_names;
|
||||
valid_names.reserve(map.size());
|
||||
for (const auto& [key, value] : map) {
|
||||
valid_names.push_back(key);
|
||||
}
|
||||
std::sort(valid_names.begin(), valid_names.end());
|
||||
|
||||
// Construct the error message
|
||||
std::ostringstream message;
|
||||
message << "Invalid name '" << name << "'. Valid names: [";
|
||||
int i = 0;
|
||||
for (const auto& key : valid_names) {
|
||||
message << "'" << key << "'";
|
||||
if (i < map.size() - 1) message << ", ";
|
||||
i++;
|
||||
}
|
||||
message << "]";
|
||||
return message.str();
|
||||
}
|
||||
|
||||
std::string IndexErrorMessage(int index, int size) {
|
||||
std::ostringstream message;
|
||||
message << "Invalid index " << index << ". Valid indices from 0 to "
|
||||
<< size - 1;
|
||||
return message.str();
|
||||
}
|
||||
} // namespace mujoco::python
|
||||
|
||||
@@ -185,6 +185,12 @@ class MjDataIndexer {
|
||||
MJDATA_VIEW_GROUPS
|
||||
#undef XGROUP
|
||||
};
|
||||
|
||||
// Returns an error message when a nonexistent name is requested from an
|
||||
// indexer, which includes all valid names.
|
||||
std::string KeyErrorMessage(const NameToIDMap& map, std::string_view name);
|
||||
// Returns an error message when an invalid numeric index is provided.
|
||||
std::string IndexErrorMessage(int index, int size);
|
||||
} // namespace mujoco::python
|
||||
|
||||
#endif // MUJOCO_PYTHON_INDEXERS_H_
|
||||
|
||||
@@ -1479,7 +1479,8 @@ This is useful for example when the MJB is not available as a file on disk.)"));
|
||||
#field, [](MjModelWrapper& m, std::string_view name) -> auto& { \
|
||||
return m.indexer().field##_by_name(name); \
|
||||
}, \
|
||||
py::return_value_policy::reference_internal);
|
||||
py::return_value_policy::reference_internal, py::arg_v("name", ""));
|
||||
|
||||
|
||||
MJMODEL_VIEW_GROUPS
|
||||
#undef XGROUP
|
||||
@@ -1493,7 +1494,7 @@ This is useful for example when the MJB is not available as a file on disk.)"));
|
||||
#altname, [](MjModelWrapper& m, std::string_view name) -> auto& { \
|
||||
return m.indexer().field##_by_name(name); \
|
||||
}, \
|
||||
py::return_value_policy::reference_internal);
|
||||
py::return_value_policy::reference_internal, py::arg_v("name", ""));
|
||||
|
||||
MJMODEL_VIEW_GROUPS_ALTNAMES
|
||||
#undef XGROUP
|
||||
@@ -1752,7 +1753,7 @@ This is useful for example when the MJB is not available as a file on disk.)"));
|
||||
#field, [](MjDataWrapper& d, std::string_view name) -> auto& { \
|
||||
return d.indexer().field##_by_name(name); \
|
||||
}, \
|
||||
py::return_value_policy::reference_internal);
|
||||
py::return_value_policy::reference_internal, py::arg_v("name", ""));
|
||||
|
||||
MJDATA_VIEW_GROUPS
|
||||
#undef XGROUP
|
||||
@@ -1766,7 +1767,7 @@ This is useful for example when the MJB is not available as a file on disk.)"));
|
||||
#altname, [](MjDataWrapper& d, std::string_view name) -> auto& { \
|
||||
return d.indexer().field##_by_name(name); \
|
||||
}, \
|
||||
py::return_value_policy::reference_internal);
|
||||
py::return_value_policy::reference_internal, py::arg_v("name", ""));
|
||||
|
||||
MJDATA_VIEW_GROUPS_ALTNAMES
|
||||
#undef XGROUP
|
||||
|
||||
Reference in New Issue
Block a user