Fix incorrect data types in Python bindings of certain arrays.

The mismatch between the `py::array_t` template argument and the type of the MuJoCo struct member being wrapped caused the `py::array_t` to be constructed via an unintended constructor overload, which in turn triggered `use-after-poison` under asan.

This bug affected the `geom`, `flex`, `elem`, and `vert` arrays in `mjContact`, and all array members in `mjrContext`.

PiperOrigin-RevId: 608632860
Change-Id: Ic227f691041b004a0fcf0cb6fa3d62c44861d9aa
This commit is contained in:
Saran Tunyasuvunakool
2024-02-20 09:25:26 -08:00
committed by Copybara-Service
parent 7c534a066c
commit abf6d41b7f
4 changed files with 72 additions and 13 deletions
+8
View File
@@ -529,6 +529,7 @@ class MuJoCoBindingsTest(parameterized.TestCase):
expected_ncon = 4
mujoco.mj_forward(self.model, self.data)
self.assertLen(self.data.contact, expected_ncon)
np.testing.assert_array_equal(self.data.contact.geom, [[0, 1]] * 4)
expected_pos = []
for contact in self.data.contact:
@@ -551,6 +552,13 @@ class MuJoCoBindingsTest(parameterized.TestCase):
self.assertLen(expected_H, expected_ncon)
np.testing.assert_array_equal(self.data.contact.H, expected_H)
expected_geom = []
for i, contact in enumerate(self.data.contact):
expected_geom.append([i, i + 1])
contact.geom = expected_geom[-1]
self.assertLen(expected_geom, expected_ncon)
np.testing.assert_array_equal(self.data.contact.geom, expected_geom)
def test_realloc_con_efc(self):
self.assertEmpty(self.data.contact)
+6 -1
View File
@@ -14,6 +14,8 @@
#include <array>
#include <cstdint>
#include <optional>
#include <type_traits>
#include <Eigen/Core>
#include <mujoco/mjrender.h>
@@ -40,7 +42,10 @@ class MjWrapper<raw::MjrContext> : public WrapperBase<raw::MjrContext> {
void Free();
#define X(var) py_array_or_tuple_t<mjtNum> var
#define X(var) \
py_array_or_tuple_t< \
std::remove_all_extents_t<decltype(raw::MjrContext::var)>> \
var
X(fogRGBA);
X(auxWidth);
X(auxHeight);
+52 -12
View File
@@ -17,12 +17,17 @@
#include <array>
#include <cstddef>
#include <functional>
#include <istream>
#include <memory>
#include <optional>
#include <ostream>
#include <sstream>
#include <string_view>
#include <unordered_map>
#include <string>
#include <type_traits>
#include <utility>
#include <vector>
#include <absl/types/span.h>
@@ -36,6 +41,7 @@
namespace mujoco::python {
namespace _impl {
template <typename T>
class WrapperBase {
public:
@@ -44,7 +50,7 @@ class WrapperBase {
T* get() { return ptr_; }
const T* get() const { return ptr_; }
const pybind11::handle owner() const { return owner_; }
pybind11::handle owner() const { return owner_; }
protected:
static void DefaultCapsuleDestructor(PyObject* pyobj) {
@@ -167,7 +173,10 @@ class MjWrapper<raw::MjOption> : public WrapperBase<raw::MjOption> {
MjWrapper(raw::MjOption* ptr, pybind11::handle owner);
~MjWrapper() = default;
#define X(var, dim) py_array_or_tuple_t<mjtNum> var;
#define X(var, dim) \
py_array_or_tuple_t< \
std::remove_all_extents_t<decltype(raw::MjOption::var)>> \
var;
MJOPTION_VECTORS
#undef X
};
@@ -188,7 +197,10 @@ class MjWrapper<raw::MjVisualHeadlight>
MjWrapper(raw::MjVisualHeadlight* ptr, pybind11::handle owner);
~MjWrapper() = default;
#define X(var) py_array_or_tuple_t<float> var
#define X(var) \
py_array_or_tuple_t< \
std::remove_all_extents_t<decltype(raw::MjVisualHeadlight::var)>> \
var
X(ambient);
X(diffuse);
X(specular);
@@ -209,7 +221,10 @@ class MjWrapper<raw::MjVisualRgba> : public WrapperBase<raw::MjVisualRgba> {
MjWrapper(raw::MjVisualRgba* ptr, pybind11::handle owner);
~MjWrapper() = default;
#define X(var) py_array_or_tuple_t<float> var
#define X(var) \
py_array_or_tuple_t< \
std::remove_all_extents_t<decltype(raw::MjVisualRgba::var)>> \
var
X(fog);
X(haze);
X(force);
@@ -271,7 +286,10 @@ class MjWrapper<raw::MjStatistic> : public WrapperBase<raw::MjStatistic> {
MjWrapper(raw::MjStatistic* ptr, pybind11::handle owner);
~MjWrapper() = default;
#define X(var) py_array_or_tuple_t<mjtNum> var
#define X(var) \
py_array_or_tuple_t< \
std::remove_all_extents_t<decltype(raw::MjStatistic::var)>> \
var
X(center);
#undef X
};
@@ -504,7 +522,10 @@ class MjWrapper<raw::MjContact> : public WrapperBase<raw::MjContact> {
MjWrapper(raw::MjContact* ptr, pybind11::handle owner);
~MjWrapper() = default;
#define X(var) py_array_or_tuple_t<mjtNum> var
#define X(var) \
py_array_or_tuple_t< \
std::remove_all_extents_t<decltype(raw::MjContact::var)>> \
var
X(pos);
X(frame);
X(friction);
@@ -628,7 +649,10 @@ class MjWrapper<raw::MjvPerturb> : public WrapperBase<raw::MjvPerturb> {
MjWrapper(MjWrapper&&) = default;
~MjWrapper() = default;
#define X(var) py_array_or_tuple_t<mjtNum> var
#define X(var) \
py_array_or_tuple_t< \
std::remove_all_extents_t<decltype(raw::MjvPerturb::var)>> \
var
X(refpos);
X(refquat);
X(refselpos);
@@ -650,7 +674,10 @@ class MjWrapper<raw::MjvCamera> : public WrapperBase<raw::MjvCamera> {
MjWrapper(MjWrapper&&) = default;
~MjWrapper() = default;
#define X(var) py_array_or_tuple_t<mjtNum> var
#define X(var) \
py_array_or_tuple_t< \
std::remove_all_extents_t<decltype(raw::MjvCamera::var)>> \
var
X(lookat);
#undef X
};
@@ -671,7 +698,10 @@ class MjWrapper<raw::MjvGLCamera> : public WrapperBase<raw::MjvGLCamera> {
explicit MjWrapper(raw::MjvGLCamera&& other);
~MjWrapper() = default;
#define X(var) py_array_or_tuple_t<float> var
#define X(var) \
py_array_or_tuple_t< \
std::remove_all_extents_t<decltype(raw::MjvGLCamera::var)>> \
var
X(pos);
X(forward);
X(up);
@@ -693,7 +723,10 @@ class MjWrapper<raw::MjvGeom> : public WrapperBase<raw::MjvGeom> {
MjWrapper(raw::MjvGeom* ptr, pybind11::handle owner);
~MjWrapper() = default;
#define X(var) py_array_or_tuple_t<float> var
#define X(var) \
py_array_or_tuple_t< \
std::remove_all_extents_t<decltype(raw::MjvGeom::var)>> \
var
X(texrepeat);
X(size);
X(pos);
@@ -717,7 +750,10 @@ class MjWrapper<raw::MjvLight> : public WrapperBase<raw::MjvLight> {
MjWrapper(raw::MjvLight* ptr, pybind11::handle owner);
~MjWrapper() = default;
#define X(var) py_array_or_tuple_t<float> var
#define X(var) \
py_array_or_tuple_t< \
std::remove_all_extents_t<decltype(raw::MjvLight::var)>> \
var
X(pos);
X(dir);
X(attenuation);
@@ -741,7 +777,10 @@ class MjWrapper<raw::MjvOption> : public WrapperBase<raw::MjvOption> {
MjWrapper(MjWrapper&&) = default;
~MjWrapper() = default;
#define X(var) py_array_or_tuple_t<mjtByte> var
#define X(var) \
py_array_or_tuple_t< \
std::remove_all_extents_t<decltype(raw::MjvOption::var)>> \
var
X(geomgroup);
X(sitegroup);
X(jointgroup);
@@ -862,6 +901,7 @@ class ScopedMsanDisabler {
void* shadow_;
};
#endif
} // namespace _impl
template <typename T>