From f1af86cf8241633bda5f27bfd1f15c08288aadcd Mon Sep 17 00:00:00 2001 From: Saran Tunyasuvunakool Date: Mon, 30 Jan 2023 05:10:39 -0800 Subject: [PATCH] Fix msan failure in `__repr__` when upgraded to NumPy 1.24. PiperOrigin-RevId: 505661236 Change-Id: I37ea339679870d23886e3c95f51451349fa1de2c --- python/mujoco/indexers.h | 2 ++ python/mujoco/structs.cc | 4 ++-- python/mujoco/structs.h | 37 +++++++++++++++++++++++++++++++++++++ 3 files changed, 41 insertions(+), 2 deletions(-) diff --git a/python/mujoco/indexers.h b/python/mujoco/indexers.h index 798bda15..4b83185f 100644 --- a/python/mujoco/indexers.h +++ b/python/mujoco/indexers.h @@ -102,6 +102,7 @@ class MjModelGroupedViewsBase { int index() const { return index_; } const std::string& name() const { return name_; } + raw::MjModel* ptr() { return m_; } protected: int index_; @@ -168,6 +169,7 @@ class MjDataGroupedViewsBase { int index() const { return index_; } const std::string& name() const { return name_; } + raw::MjData* ptr() { return d_; } protected: int index_; diff --git a/python/mujoco/structs.cc b/python/mujoco/structs.cc index 73612b57..9b1d28e7 100644 --- a/python/mujoco/structs.cc +++ b/python/mujoco/structs.cc @@ -1551,7 +1551,7 @@ This is useful for example when the MJB is not available as a file on disk.)")); using GroupedViews = MjModelGroupedViews; \ py::class_ groupedViews(m, "_" #MjModelGroupedViews); \ FIELD_XMACROS \ - groupedViews.def("__repr__", StructRepr); \ + groupedViews.def("__repr__", MjModelStructRepr); \ groupedViews.def_property_readonly("id", [](GroupedViews& views) { \ return views.index(); \ }); \ @@ -1868,7 +1868,7 @@ This is useful for example when the MJB is not available as a file on disk.)")); using GroupedViews = MjDataGroupedViews; \ py::class_ groupedViews(m, "_" #MjDataGroupedViews); \ FIELD_XMACROS \ - groupedViews.def("__repr__", StructRepr); \ + groupedViews.def("__repr__", MjDataStructRepr); \ groupedViews.def_property_readonly("id", [](GroupedViews& views) { \ return views.index(); \ }); \ diff --git a/python/mujoco/structs.h b/python/mujoco/structs.h index 8d282619..3e80f1e7 100644 --- a/python/mujoco/structs.h +++ b/python/mujoco/structs.h @@ -808,6 +808,25 @@ using MjvFigureWrapper = MjWrapper; template <> struct enable_if_mj_struct { using type = void; }; + +#ifdef MEMORY_SANITIZER +template +class ScopedMsanDisabler { + public: + ScopedMsanDisabler(T* mj) + : mj_(mj), shadow_(std::malloc(mj_->nbuffer)) { + __msan_copy_shadow(shadow_, mj_->buffer, mj_->nbuffer); + __msan_unpoison(mj_->buffer, mj_->nbuffer); + } + ~ScopedMsanDisabler() { + __msan_copy_shadow(mj_->buffer, shadow_, mj_->nbuffer); + std::free(shadow_); + } + private: + T* mj_; + void* shadow_; +}; +#endif } // namespace _impl template @@ -1078,6 +1097,24 @@ std::string StructRepr(pybind11::object self) { return result.str(); } +template +std::string MjModelStructRepr(pybind11::object self) { +#ifdef MEMORY_SANITIZER + _impl::ScopedMsanDisabler msan_disabler( + pybind11::cast(self).ptr()); +#endif + return StructRepr(self); +} + +template +std::string MjDataStructRepr(pybind11::object self) { +#ifdef MEMORY_SANITIZER + _impl::ScopedMsanDisabler msan_disabler( + pybind11::cast(self).ptr()); +#endif + return StructRepr(self); +} + template void DefineStructFunctions(pybind11::class_ c) { c.def("__eq__", StructsEqual);