diff --git a/doc/APIreference/functions.rst b/doc/APIreference/functions.rst index 0713d163..f7dd107f 100644 --- a/doc/APIreference/functions.rst +++ b/doc/APIreference/functions.rst @@ -2161,7 +2161,25 @@ Get compiler timing diagnostics from spec, returns pointer to array of size mjNC .. mujoco-include:: mjs_isWarning -Return 1 if compiler error is a warning. +Return 1 if compiler error is a warning. Deprecated: use mjs_numWarnings(s) > 0. + +.. _mjs_numWarnings: + +`mjs_numWarnings <#mjs_numWarnings>`__ +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +.. mujoco-include:: mjs_numWarnings + +Get number of warnings accumulated in the spec. + +.. _mjs_getWarning: + +`mjs_getWarning <#mjs_getWarning>`__ +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +.. mujoco-include:: mjs_getWarning + +Get the i-th warning message (returns nullptr if index out of bounds). .. _Miscellaneous: diff --git a/doc/changelog.rst b/doc/changelog.rst index dd5c5de6..4beda810 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -20,6 +20,8 @@ General - New types: :ref:`mjtLogLevel`, :ref:`mjtLogTopic`, :ref:`mjLogMessage`, :ref:`mjLogConfig`. - The legacy callbacks :ref:`mju_user_error` and :ref:`mju_user_warning` are deprecated but remain functional. +- Added :ref:`mjs_numWarnings` and :ref:`mjs_getWarning` for retrieving all warnings accumulated during model + compilation and attachment. Deprecated :ref:`mjs_isWarning` in favor of ``mjs_numWarnings(s) > 0``. - Improved primal solver convergence under float32. Improvements initially proposed by :github:user:`n3b` in :issue:`2313` and :github:user:`denzeler-nvidia` in :doc:`MJWarp ` pull request `1374 `__. diff --git a/doc/includes/references.h b/doc/includes/references.h index b47fa5cb..66c834c4 100644 --- a/doc/includes/references.h +++ b/doc/includes/references.h @@ -3555,6 +3555,8 @@ void mju_writeLog(const char* type, const char* msg); const char* mjs_getError(mjSpec* s); const double* mjs_getTimer(mjSpec* s); int mjs_isWarning(mjSpec* s); +int mjs_numWarnings(const mjSpec* spec); +const char* mjs_getWarning(const mjSpec* spec, int index); void mju_zero3(mjtNum res[3]); void mju_copy3(mjtNum res[3], const mjtNum data[3]); void mju_scl3(mjtNum res[3], const mjtNum vec[3], mjtNum scl); diff --git a/include/mujoco/mujoco.h b/include/mujoco/mujoco.h index 63b35d4f..693bee10 100644 --- a/include/mujoco/mujoco.h +++ b/include/mujoco/mujoco.h @@ -1015,9 +1015,14 @@ MJAPI const char* mjs_getError(mjSpec* s); // Get compiler timing diagnostics from spec, returns pointer to array of size mjNCTIMER. MJAPI const double* mjs_getTimer(mjSpec* s); -// Return 1 if compiler error is a warning. +// Return 1 if compiler error is a warning. Deprecated: use mjs_numWarnings(s) > 0. MJAPI int mjs_isWarning(mjSpec* s); +// Get number of warnings accumulated in the spec. +MJAPI int mjs_numWarnings(const mjSpec* spec); + +// Get the i-th warning message (returns nullptr if index out of bounds). +MJAPI const char* mjs_getWarning(const mjSpec* spec, int index); //---------------------------------- Standard math ------------------------------------------------- diff --git a/python/mujoco/introspect/functions.py b/python/mujoco/introspect/functions.py index 8a0c2e7b..fc817cf9 100644 --- a/python/mujoco/introspect/functions.py +++ b/python/mujoco/introspect/functions.py @@ -6486,7 +6486,41 @@ FUNCTIONS: Mapping[str, FunctionDecl] = dict([ ), ), ), - doc='Return 1 if compiler error is a warning.', + doc='Return 1 if compiler error is a warning. Deprecated: use mjs_numWarnings(s) > 0.', # pylint: disable=line-too-long + )), + ('mjs_numWarnings', + FunctionDecl( + name='mjs_numWarnings', + return_type=ValueType(name='int'), + parameters=( + FunctionParameterDecl( + name='spec', + type=PointerType( + inner_type=ValueType(name='mjSpec', is_const=True), + ), + ), + ), + doc='Get number of warnings accumulated in the spec.', + )), + ('mjs_getWarning', + FunctionDecl( + name='mjs_getWarning', + return_type=PointerType( + inner_type=ValueType(name='char', is_const=True), + ), + parameters=( + FunctionParameterDecl( + name='spec', + type=PointerType( + inner_type=ValueType(name='mjSpec', is_const=True), + ), + ), + FunctionParameterDecl( + name='index', + type=ValueType(name='int'), + ), + ), + doc='Get the i-th warning message (returns nullptr if index out of bounds).', # pylint: disable=line-too-long )), ('mju_zero3', FunctionDecl( diff --git a/python/mujoco/specs_test.py b/python/mujoco/specs_test.py index 17b803c8..282a0482 100644 --- a/python/mujoco/specs_test.py +++ b/python/mujoco/specs_test.py @@ -597,6 +597,18 @@ class SpecsTest(absltest.TestCase): with self.assertRaisesRegex(ValueError, expected_error): spec.compile() + def test_compile_warnings(self): + xml = """ + + + + + + """ + spec = mujoco.MjSpec.from_string(xml) + with self.assertWarnsRegex(UserWarning, 'is not rigid'): + spec.compile() + def test_recompile(self): # Create a spec. spec = mujoco.MjSpec() diff --git a/python/mujoco/specs_wrapper.cc b/python/mujoco/specs_wrapper.cc index 31d52e41..049199e1 100644 --- a/python/mujoco/specs_wrapper.cc +++ b/python/mujoco/specs_wrapper.cc @@ -24,6 +24,7 @@ #include #include "errors.h" #include "indexers.h" // IWYU pragma: keep +#include "private.h" #include "raw.h" #include "structs.h" // IWYU pragma: keep #include @@ -118,7 +119,14 @@ raw::MjModel* MjSpec::Compile(mjVFS* vfs) { raw::MjModel* m; { py::gil_scoped_release no_gil; + + // Install a no-op handler to suppress stderr output from warnings. + // Compile() installs its own (setjmp/longjmp) handler and then chains to + // prev. We want to raise a `warnings.warn`, so pass in a no-op handler. + mjfLogHandler prev = + _mjPRIVATE_setTlsLogHandler([](const mjLogMessage*) {}); m = mj_compile(ptr, vfs); + _mjPRIVATE_setTlsLogHandler(prev); } if (local_vfs.has_value()) { @@ -127,9 +135,17 @@ raw::MjModel* MjSpec::Compile(mjVFS* vfs) { local_vfs = std::nullopt; } - if (!m || mjs_isWarning(ptr)) { + if (!m) { throw py::value_error(mjs_getError(ptr)); } + + int num_warnings = mjs_numWarnings(ptr); + if (num_warnings > 0) { + py::object warnings = py::module_::import("warnings"); + for (int i = 0; i < num_warnings; ++i) { + warnings.attr("warn")(mjs_getWarning(ptr, i)); + } + } return m; } diff --git a/src/user/user_api.cc b/src/user/user_api.cc index ad4f4c36..432be2f7 100644 --- a/src/user/user_api.cc +++ b/src/user/user_api.cc @@ -483,15 +483,37 @@ const double* mjs_getTimer(mjSpec* s) { return modelC->timer; } - - -// check if model has warnings +// check if model has warnings (but no error) +// TODO(tassa): delete this function int mjs_isWarning(mjSpec* s) { + if (!s) { + return 0; + } mjCModel* modelC = static_cast(s->element); - return modelC->GetError().warning; + return modelC->GetError().message[0] == '\0' && + !modelC->GetWarnings().empty(); } +// get number of warnings +int mjs_numWarnings(const mjSpec* spec) { + if (!spec) { + return 0; + } + const mjCModel* modelC = static_cast(spec->element); + return static_cast(modelC->GetWarnings().size()); +} +// get the i-th warning message +const char* mjs_getWarning(const mjSpec* spec, int index) { + if (!spec) { + return nullptr; + } + const mjCModel* modelC = static_cast(spec->element); + if (index < 0 || index >= static_cast(modelC->GetWarnings().size())) { + return nullptr; + } + return modelC->GetWarnings()[index].c_str(); +} // delete model void mj_deleteSpec(mjSpec* s) { diff --git a/src/user/user_mesh.cc b/src/user/user_mesh.cc index d74d97b9..ece381b4 100644 --- a/src/user/user_mesh.cc +++ b/src/user/user_mesh.cc @@ -1546,7 +1546,7 @@ void mjCMesh::Process() { for (int i = 0; i < nface(); i++) { SetBoundingVolume(i, dvert.data()); } - tree_.CreateBVH(); + tree_.CreateBVH(model, this); } mesh_timer_[mjCTIMER_MESH_BVH] += Seconds(Clock::now() - t0).count(); @@ -5451,7 +5451,7 @@ void mjCFlex::CreateBVH() { // create hierarchy tree.RemoveInactiveVolumes(nbvh); - tree.CreateBVH(); + tree.CreateBVH(model, this); } diff --git a/src/user/user_model.cc b/src/user/user_model.cc index f10cd4fd..10eb4db2 100644 --- a/src/user/user_model.cc +++ b/src/user/user_model.cc @@ -1239,6 +1239,7 @@ void mjCModel::Clear() { hasImplicitPluginElem = false; compiled = false; errInfo = mjCError(); + ClearCompileWarnings(); qpos0.clear(); } @@ -1508,7 +1509,21 @@ const mjCError& mjCModel::GetError() const { return errInfo; } +// add warning to vector (immediate delivery outside compile) +void mjCModel::AddWarning(std::string msg, const mjCBase* obj) { + if (obj) { + msg += "\nElement name '" + obj->name + "', id " + std::to_string(obj->id); + if (!obj->info.empty()) { + msg += ", " + obj->info; + } + } + // outside compile: deliver immediately via normal handler chain + if (!compiling_) { + mju_warning("%s", msg.c_str()); + } + warnings_.push_back(std::move(msg)); +} // pointer to world body mjCBody* mjCModel::GetWorld() { @@ -3551,8 +3566,10 @@ void mjCModel::CopyObjects(mjModel* m) { if (!pfl->rigid && m->flex_edgeequality[i] == 0 && !pfl->edgestiffness && !pfl->edgedamping && !pfl->damping && pfl->bending.empty()) { - mju_warning("flex '%s' is not rigid and has no equality constraints " - "or passive forces", pfl->name.c_str()); + AddWarning("flex '" + pfl->name + + "' is not rigid and has no equality constraints or " + "passive forces", + pfl); } // copy bvh data (flex aabb computed dynamically in mjData) @@ -4212,9 +4229,10 @@ template void mjCModel::RestoreState( // resolve keyframe references void mjCModel::StoreKeyframes(mjCModel* dest) { if (this != dest && !key_pending_.empty()) { - mju_warning( - "Child model has pending keyframes. They will not be namespaced correctly. " - "To prevent this, compile the child model before attaching it again."); + dest->AddWarning( + "Child model has pending keyframes. They will not be namespaced " + "correctly. " + "To prevent this, compile the child model before attaching it again."); } // do not change compilation quantities in case the user wants to recompile preserving the state @@ -4633,10 +4651,17 @@ static void compilerLogHandler(const mjLogMessage* msg) { mju::strcpy_arr(errortext, msg->subject); std::longjmp(error_jmp_buf, 1); } else if (msg->level == mjLOG_WARNING) { + // buffer for structured capture (append, not overwrite) if (local_warningtext_ptr) { - *local_warningtext_ptr = msg->subject; + if (!local_warningtext_ptr->empty()) { + *local_warningtext_ptr += '\n'; + } + *local_warningtext_ptr += msg->subject; } else { - mju::strcpy_arr(warningtext, msg->subject); + if (warningtext[0]) { + mju::strcat_arr(warningtext, "\n"); + } + mju::strcat_arr(warningtext, msg->subject); } } } @@ -4661,10 +4686,15 @@ mjModel* mjCModel::Compile(const mjVFS* vfs, mjModel** m) { mjModel* volatile model = (m && *m) ? *m : nullptr; mjData* volatile data = nullptr; - // save log handler - mjfLogHandler save_handler = _mjPRIVATE_setTlsLogHandler(compilerLogHandler); + // install compiler log handler (captures warnings silently) + mjfLogHandler prev_tls = _mjPRIVATE_setTlsLogHandler(compilerLogHandler); errInfo = mjCError(); + + // set flag so warnings are captured in the spec vector rather than delivered + // immediately + compiling_ = true; + ClearCompileWarnings(); warningtext[0] = 0; try { @@ -4677,7 +4707,7 @@ mjModel* mjCModel::Compile(const mjVFS* vfs, mjModel** m) { // also include the last warning that was issued. this is useful for // warnings that came out of plugin implementations. if (warningtext[0]) { - error_msg += "\n"; + error_msg += '\n'; error_msg += warningtext; } throw mjCError(0, "engine error: %s", error_msg.c_str()); @@ -4701,13 +4731,21 @@ mjModel* mjCModel::Compile(const mjVFS* vfs, mjModel** m) { } // restore handler, return 0 - _mjPRIVATE_setTlsLogHandler(save_handler); + _mjPRIVATE_setTlsLogHandler(prev_tls); + compiling_ = false; return nullptr; } - // restore log handler, mark as compiled, return mjModel - _mjPRIVATE_setTlsLogHandler(save_handler); + // restore log handler + _mjPRIVATE_setTlsLogHandler(prev_tls); + compiling_ = false; compiled = true; + + // play back compile warnings through the normal handler chain + for (int i = num_attach_warnings_; i < warnings_.size(); ++i) { + mju_warning("%s", warnings_[i].c_str()); + } + return model; } @@ -5353,7 +5391,22 @@ void mjCModel::TryCompile(mjModel*& m, mjData*& d, const mjVFS* vfs) { throw mjCError(0, "could not create mjData"); } + // pass compiler warnings into structured warning vector before validation + if (warningtext[0]) { + std::string warnings(warningtext); + std::istringstream stream(warnings); + std::string line; + while (std::getline(stream, line)) { + if (!line.empty()) { + AddWarning(line); + } + } + } + // test forward simulation unless asleep_init is true (potentially expensive) + // reset warningtext: engine warnings from validation are not compiler + // warnings + warningtext[0] = 0; if (!asleep_init) { mj_step(m, d); } @@ -5364,11 +5417,6 @@ void mjCModel::TryCompile(mjModel*& m, mjData*& d, const mjVFS* vfs) { m->opt.enableflags = enableflags; d = nullptr; - // pass warning back - if (warningtext[0]) { - mju::strcpy_arr(errInfo.message, warningtext); - errInfo.warning = true; - } // save signature m->signature = Signature(); diff --git a/src/user/user_model.h b/src/user/user_model.h index 3af7167e..b7dc3e2b 100644 --- a/src/user/user_model.h +++ b/src/user/user_model.h @@ -249,6 +249,23 @@ class mjCModel : public mjCModel_, private mjSpec { bool IsCompiled() const; // is model already compiled const mjCError& GetError() const; // get reference of error object void SetError(const mjCError& error) { errInfo = error; } // set value of error object + void AddWarning(std::string msg, // add warning to vector + const mjCBase* obj = nullptr); + const std::vector& GetWarnings() + const { // get accumulated warnings + return warnings_; + } + void ClearWarnings() { + warnings_.clear(); + num_attach_warnings_ = 0; + } // clear all warnings + void ClearCompileWarnings() { + warnings_.resize(num_attach_warnings_); + } // clear compile warnings + void SetAttachWarningBoundary() { // snapshot attach warning count + num_attach_warnings_ = warnings_.size(); + } + mjCBody* GetWorld(); // pointer to world body mjCDef* FindDefault(const std::string& name) const; // find defaults class name mjCDef* AddDefault(std::string name, mjCDef* parent = nullptr); // add defaults class to array @@ -488,10 +505,15 @@ class mjCModel : public mjCModel_, private mjSpec { // expand all keyframes in the model void ExpandAllKeyframes(); - mjListKeyMap ids; // map from object names to ids - mjCError errInfo; // last error info + mjListKeyMap ids; // map from object names to ids + mjCError errInfo; // last error info + std::vector + warnings_; // chronological list of non-fatal warnings + int num_attach_warnings_ = + 0; // boundary: [0, n) are attach, [n, size) are compile + bool compiling_ = false; // true during Compile() std::vector key_pending_; // attached keyframes - bool deepcopy_; // copy objects when attaching + bool deepcopy_; // copy objects when attaching bool attached_ = false; // true if model is attached to a parent model std::unordered_map compiler2spec_; // map from compiler to spec std::vector detached_; // list of detached objects diff --git a/src/user/user_objects.cc b/src/user/user_objects.cc index b4d76a10..297775a8 100644 --- a/src/user/user_objects.cc +++ b/src/user/user_objects.cc @@ -38,12 +38,11 @@ #include #include -#include "lodepng.h" -#include "cc/array_safety.h" -#include "engine/engine_passive.h" -#include "engine/engine_support.h" +#include "lodepng.h" // NOLINT #include #include +#include "cc/array_safety.h" +#include "engine/engine_passive.h" #include "user/user_api.h" #include "user/user_cache.h" #include "user/user_model.h" @@ -197,7 +196,6 @@ mjCError::mjCError(const mjCBase* obj, const char* msg, const char* str, int pos char temp[600]; // init - warning = false; if (obj || msg) { mju::sprintf_arr(message, "Error"); } else { @@ -396,13 +394,13 @@ mjCBoundingVolumeHierarchy::AddBoundingVolume(const int* id, int contype, int co // create bounding volume hierarchy -void mjCBoundingVolumeHierarchy::CreateBVH() { +void mjCBoundingVolumeHierarchy::CreateBVH(mjCModel* model, + const mjCBase* owner) { std::vector elements; Make(elements); - MakeBVH(elements.begin(), elements.end()); + MakeBVH(elements.begin(), elements.end(), 0, model, owner); } - void mjCBoundingVolumeHierarchy::Make(std::vector& elements) { // precompute the positions of each element in the hierarchy's axes, and drop // visual-only elements. @@ -424,8 +422,9 @@ void mjCBoundingVolumeHierarchy::Make(std::vector& elements) { // compute bounding volume hierarchy int mjCBoundingVolumeHierarchy::MakeBVH( - std::vector::iterator elements_begin, - std::vector::iterator elements_end, int lev) { + std::vector::iterator elements_begin, + std::vector::iterator elements_end, int lev, mjCModel* model, + const mjCBase* owner) { int nelements = elements_end - elements_begin; if (nelements == 0) { return -1; @@ -525,11 +524,13 @@ int mjCBoundingVolumeHierarchy::MakeBVH( // recursive calls if (m > 0) { - child_[2*index + 0] = MakeBVH(elements_begin, elements_begin + m, lev + 1); + child_[2 * index + 0] = + MakeBVH(elements_begin, elements_begin + m, lev + 1, model, owner); } if (m != nelements) { - child_[2*index + 1] = MakeBVH(elements_begin + m, elements_end, lev + 1); + child_[2 * index + 1] = + MakeBVH(elements_begin + m, elements_end, lev + 1, model, owner); } // SHOULD NOT OCCUR @@ -539,14 +540,12 @@ int mjCBoundingVolumeHierarchy::MakeBVH( } if (lev > mjMAXTREEDEPTH) { - mju_warning("max tree depth exceeded in body=%s", name_.c_str()); + model->AddWarning("max tree depth exceeded", owner); } return index; } - - //------------------------- class mjCOctree implementation -------------------------------------------- void mjCOctree::CopyLevel(int* level) const { @@ -2631,7 +2630,7 @@ void mjCBody::ComputeBVH() { tree.AddBoundingVolume(&geom->id, geom->contype, geom->conaffinity, geom->pos, geom->quat, geom->aabb); } - tree.CreateBVH(); + tree.CreateBVH(model, this); } @@ -3871,7 +3870,7 @@ void mjCGeom::SetFluidCoefs(void) { // compute bounding box void mjCGeom::ComputeAABB(void) { - double aamm[6]; // axis-aligned bounding box in (min, max) format + double aamm[6]; // axis-aligned bounding box in (min, max) format switch (type) { case mjGEOM_HFIELD: aamm[0] = -hfield->size[0]; diff --git a/src/user/user_objects.h b/src/user/user_objects.h index ed93cf48..ed40733c 100644 --- a/src/user/user_objects.h +++ b/src/user/user_objects.h @@ -83,7 +83,6 @@ class [[nodiscard]] mjCError { int pos2 = 0); char message[500]; // error message - bool warning; // is this a warning instead of error }; // alternative specifications of frame orientation @@ -172,7 +171,7 @@ struct mjCBoundingVolumeHierarchy_ { class mjCBoundingVolumeHierarchy : public mjCBoundingVolumeHierarchy_ { public: // make bounding volume hierarchy - void CreateBVH(); + void CreateBVH(mjCModel* model, const mjCBase* owner); void Set(double ipos_element[3], double iquat_element[4]); void AllocateBoundingVolumes(int nleaf); void RemoveInactiveVolumes(int nmax); @@ -210,7 +209,8 @@ class mjCBoundingVolumeHierarchy : public mjCBoundingVolumeHierarchy_ { }; void Make(std::vector& elements); int MakeBVH(std::vector::iterator elements_begin, - std::vector::iterator elements_end, int lev = 0); + std::vector::iterator elements_end, int lev, + mjCModel* model, const mjCBase* owner); }; diff --git a/src/xml/xml_api.cc b/src/xml/xml_api.cc index 5fab0475..c74a61cb 100644 --- a/src/xml/xml_api.cc +++ b/src/xml/xml_api.cc @@ -58,8 +58,16 @@ mjModel* mj_loadXML(const char* filename, const mjVFS* vfs, } // handle compile warning - if (mjs_isWarning(spec.get())) { - mjCopyError(error, mjs_getError(spec.get()), error_sz); + int num_warnings = mjs_numWarnings(spec.get()); + if (num_warnings > 0) { + std::string all_warnings; + for (int i = 0; i < num_warnings; ++i) { + if (!all_warnings.empty()) { + all_warnings += '\n'; + } + all_warnings += mjs_getWarning(spec.get(), i); + } + mjCopyError(error, all_warnings.c_str(), error_sz); } else if (error) { error[0] = '\0'; } diff --git a/test/engine/engine_forward_test.cc b/test/engine/engine_forward_test.cc index 401358c1..f77ec394 100644 --- a/test/engine/engine_forward_test.cc +++ b/test/engine/engine_forward_test.cc @@ -80,7 +80,9 @@ struct ActLimitedTestCase { mjtIntegrator integrator; }; -using ParametrizedForwardTest = ::testing::TestWithParam; +class ParametrizedForwardTest + : public MujocoTest, + public ::testing::WithParamInterface {}; TEST_P(ParametrizedForwardTest, ActLimited) { static constexpr char xml[] = R"( diff --git a/test/engine/engine_sensor_test.cc b/test/engine/engine_sensor_test.cc index 1332ba99..076f1c5f 100644 --- a/test/engine/engine_sensor_test.cc +++ b/test/engine/engine_sensor_test.cc @@ -1753,6 +1753,8 @@ TEST_F(SensorTest, InsideSiteFlexBody) { )"; char error[1024] = {0}; + EXPECT_CALL(mock_warning_handler, Warn(testing::HasSubstr("is not rigid"))) + .WillOnce(testing::Return()); mjModel* m = LoadModelFromString(xml, error, sizeof(error)); ASSERT_THAT(m, NotNull()) << error; mjData* d = mj_makeData(m); @@ -1839,6 +1841,8 @@ TEST_F(SensorTest, FlexContactSensors) { )"; char error[1024] = {0}; + EXPECT_CALL(mock_warning_handler, Warn(testing::HasSubstr("is not rigid"))) + .WillOnce(testing::Return()); mjModel* m = LoadModelFromString(xml, error, sizeof(error)); ASSERT_THAT(m, NotNull()) << error; mjData* d = mj_makeData(m); diff --git a/test/engine/engine_util_container_test.cc b/test/engine/engine_util_container_test.cc index 77c65ece..35e4bb19 100644 --- a/test/engine/engine_util_container_test.cc +++ b/test/engine/engine_util_container_test.cc @@ -45,7 +45,9 @@ constexpr int GetExpectedStackUsageBytes() { } } -TEST(TestMjArrayList, TestMjArrayListSingleThreaded) { +class TestMjArrayList : public MujocoTest {}; + +TEST_F(TestMjArrayList, TestMjArrayListSingleThreaded) { std::array error; mjModel* m = LoadModelFromString("", error.data(), error.size()); ASSERT_THAT(m, NotNull()) << "Failed to load model: " << error.data(); @@ -81,7 +83,7 @@ TEST(TestMjArrayList, TestMjArrayListSingleThreaded) { mj_deleteModel(m); } -TEST(TestMjArrayList, ZeroInitialCapacity) { +TEST_F(TestMjArrayList, ZeroInitialCapacity) { char error[1024]; mjModel* m = LoadModelFromString("", error, sizeof(error)); ASSERT_THAT(m, NotNull()) << "Failed to load model: " << error; diff --git a/test/fixture.cc b/test/fixture.cc index 1f3c7d67..26d0d4d1 100644 --- a/test/fixture.cc +++ b/test/fixture.cc @@ -24,7 +24,6 @@ #include #include #include - #include #include @@ -33,6 +32,7 @@ #include #include #include +#include #include #include #include @@ -41,35 +41,103 @@ #include "src/xml/xml_global.h" namespace mujoco { +namespace { +using ::testing::_; +using ::testing::Not; +using ::testing::Return; +using ::testing::Truly; + +// Returns true if the warning matches a known benign warning to ignore. +bool IsBenignWarning(const std::string& msg) { + static const char* const kBenignWarnings[] = { + "is not rigid and has no equality constraints", + }; + for (const char* warning : kBenignWarnings) { + if (absl::StrContains(msg, warning)) { + return true; + } + } + return false; +} +} // namespace + +thread_local MockWarningHandler* MockWarningHandler::active_handler = nullptr; + +// Registers this handler as the active one. +MockWarningHandler::MockWarningHandler() { + prev_ = active_handler; + active_handler = this; + + // By default, ignore matches in the benign warnings list + ON_CALL(*this, Warn(Truly(IsBenignWarning))) + .WillByDefault([](const std::string&) {}); + + // Fail on all other warnings + ON_CALL(*this, Warn(Not(Truly(IsBenignWarning)))) + .WillByDefault([](const std::string& msg) { + ADD_FAILURE() << "mju_user_warning: " << msg; + }); +} + +// Restores the previously active warning handler. +MockWarningHandler::~MockWarningHandler() { active_handler = prev_; } + +// Configures the mock warning handler to ignore all warnings. +void MockWarningHandler::ExpectWarnings() { + EXPECT_CALL(*this, Warn(_)).WillRepeatedly(Return()); +} + +// Returns the active warning handler. +MockWarningHandler* MockWarningHandler::GetActive() { return active_handler; } + namespace { using ::testing::NotNull; ABSL_CONST_INIT static absl::Mutex handlers_mutex(absl::kConstInit); static int guard_count ABSL_GUARDED_BY(handlers_mutex) = 0; +static mjfLogHandler prev_log_handler ABSL_GUARDED_BY(handlers_mutex) = nullptr; -void default_mj_error_handler(const char* msg) { - FAIL() << "mju_user_error: " << msg; -} +void default_mj_log_handler(const mjLogMessage* msg) { + std::string subject = msg->subject; + if (msg->func) { + subject = std::string(msg->func) + ": " + msg->subject; + } -void default_mj_warning_handler(const char* msg) { - ADD_FAILURE() << "mju_user_warning: " << msg; + if (msg->level == mjLOG_ERROR) { + if (mju_user_error) { + mju_user_error(subject.c_str()); + } else { + FAIL() << "mju_user_error: " << subject; + } + } else if (msg->level == mjLOG_WARNING) { + std::string full_msg = subject; + if (msg->body) { + full_msg += "\n" + std::string(msg->body); + } + if (mju_user_warning) { + mju_user_warning(full_msg.c_str()); + } else if (auto* handler = MockWarningHandler::GetActive()) { + handler->Warn(full_msg); + } else { + ADD_FAILURE() << "mju_user_warning: " << full_msg; + } + } } } // namespace MujocoErrorTestGuard::MujocoErrorTestGuard() { absl::MutexLock lock(handlers_mutex); if (++guard_count == 1) { - mju_user_error = default_mj_error_handler; - mju_user_warning = default_mj_warning_handler; + prev_log_handler = mju_setLogHandler(default_mj_log_handler); } } MujocoErrorTestGuard::~MujocoErrorTestGuard() { absl::MutexLock lock(handlers_mutex); if (--guard_count == 0) { - mju_user_error = nullptr; - mju_user_warning = nullptr; + mju_setLogHandler(prev_log_handler); + prev_log_handler = nullptr; } } @@ -97,11 +165,29 @@ mjModel* LoadModelFromString(std::string_view xml, char* error, if (spec) { model = mj_compile(spec, vfs); - if (error && (!model || mjs_isWarning(spec))) { - strncpy(error, mjs_getError(spec), error_size); - error[error_size - 1] = '\0'; + if (error) { + if (!model) { + strncpy(error, mjs_getError(spec), error_size); + error[error_size - 1] = '\0'; + } else { + int num_warnings = mjs_numWarnings(spec); + if (num_warnings > 0) { + std::string all_warnings; + for (int i = 0; i < num_warnings; ++i) { + if (!all_warnings.empty()) { + all_warnings += '\n'; + } + all_warnings += mjs_getWarning(spec, i); + } + strncpy(error, all_warnings.c_str(), error_size); + error[error_size - 1] = '\0'; + } else { + error[0] = '\0'; + } + } } } + SetGlobalXmlSpec(spec); return model; } diff --git a/test/fixture.h b/test/fixture.h index 07e23126..805ab626 100644 --- a/test/fixture.h +++ b/test/fixture.h @@ -44,7 +44,7 @@ namespace mujoco { inline mjtNum MjTolScale() { static const mjtNum scale = []() { const char* env = std::getenv("MJTOL_SCALE"); - return env ? std::atof(env) : 1.0; + return env ? std::strtod(env, nullptr) : 1.0; }(); return scale; } @@ -106,6 +106,28 @@ class MujocoErrorTestGuard { ~MujocoErrorTestGuard(); }; +// Mock handler for capturing and verifying mju_warning logs. +class MockWarningHandler { + public: + // Constructor that registers this handler as the active one. + MockWarningHandler(); + // Destructor that restores the previously active handler. + ~MockWarningHandler(); + + // Mock method called when a warning is intercepted. + MOCK_METHOD(void, Warn, (const std::string& msg)); + + // Allow any number of warnings without triggering test failure. + void ExpectWarnings(); + + // Returns the thread-local active mock warning handler. + static MockWarningHandler* GetActive(); + + private: + static thread_local MockWarningHandler* active_handler; + MockWarningHandler* prev_ = nullptr; +}; + // A test fixture which simplifies writing tests for the MuJoCo C API. // By default, any MuJoCo operation which triggers a warning or error will // trigger a test failure. @@ -128,6 +150,9 @@ class MujocoTest : public ::testing::Test { } ~MujocoTest() { mj_freeLastXML(); } + protected: + MockWarningHandler mock_warning_handler; + private: MujocoErrorTestGuard error_guard; }; diff --git a/test/fixture_test.cc b/test/fixture_test.cc index e893024c..ca779739 100644 --- a/test/fixture_test.cc +++ b/test/fixture_test.cc @@ -34,6 +34,13 @@ TEST_F(MujocoTestTest, MjUserWarningFailsTest) { EXPECT_NONFATAL_FAILURE(mju_warning("Warning."), "Warning."); } +TEST_F(MujocoTestTest, BenignWarningDoesNotFailTest) { + // Warnings in the benign list should not trigger test failures + mju_warning( + "flex 'soft' is not rigid and has no equality constraints " + "or passive forces"); +} + TEST_F(MujocoTestTest, MjUserErrorFailsTest) { EXPECT_FATAL_FAILURE(mju_error("Error."), "Error."); } diff --git a/test/user/user_api_test.cc b/test/user/user_api_test.cc index 701381bb..cf0de3a8 100644 --- a/test/user/user_api_test.cc +++ b/test/user/user_api_test.cc @@ -3291,5 +3291,168 @@ TEST_F(MujocoTest, CompilerTimers) { mj_deleteSpec(spec); } +// -------------------- test compile warning infrastructure -------------------- + +TEST_F(MujocoTest, CompileWarningCount) { + static constexpr char xml[] = R"( + + + + + + + + + + + )"; + std::array error; + mjSpec* spec = mj_parseXMLString(xml, 0, error.data(), error.size()); + ASSERT_THAT(spec, NotNull()) << error.data(); + + mjModel* model = mj_compile(spec, 0); + ASSERT_THAT(model, NotNull()); + + // flex with no passive forces should produce a warning + EXPECT_GT(mjs_numWarnings(spec), 0); + EXPECT_THAT(mjs_getWarning(spec, 0), HasSubstr("not rigid")); + + mj_deleteModel(model); + mj_deleteSpec(spec); +} + +TEST_F(MujocoTest, CompileWarningOutOfBounds) { + mjSpec* spec = mj_makeSpec(); + mjsBody* world = mjs_findBody(spec, "world"); + mjsGeom* geom = mjs_addGeom(world, 0); + geom->size[0] = 1; + + mjModel* model = mj_compile(spec, 0); + ASSERT_THAT(model, NotNull()); + + // no warnings expected for simple model + EXPECT_EQ(mjs_numWarnings(spec), 0); + EXPECT_THAT(mjs_getWarning(spec, 0), IsNull()); + EXPECT_THAT(mjs_getWarning(spec, -1), IsNull()); + + // nullptr spec should not crash + EXPECT_EQ(mjs_numWarnings(nullptr), 0); + EXPECT_THAT(mjs_getWarning(nullptr, 0), IsNull()); + + mj_deleteModel(model); + mj_deleteSpec(spec); +} + +TEST_F(MujocoTest, RecompileClearsCompileWarnings) { + static constexpr char xml[] = R"( + + + + + + + + + + + )"; + std::array error; + mjSpec* spec = mj_parseXMLString(xml, 0, error.data(), error.size()); + ASSERT_THAT(spec, NotNull()) << error.data(); + + mjModel* model = mj_compile(spec, 0); + ASSERT_THAT(model, NotNull()); + int first_count = mjs_numWarnings(spec); + EXPECT_GT(first_count, 0); + + // recompile — warnings should be regenerated, not accumulated + mj_deleteModel(model); + model = mj_compile(spec, 0); + ASSERT_THAT(model, NotNull()); + EXPECT_EQ(mjs_numWarnings(spec), first_count); + + mj_deleteModel(model); + mj_deleteSpec(spec); +} + +TEST_F(MujocoTest, LoadXMLWarningInErrorBuffer) { + static constexpr char xml[] = R"( + + + + + + + + + + + )"; + + // write xml to VFS + mjVFS vfs; + mj_defaultVFS(&vfs); + mj_addBufferVFS(&vfs, "model.xml", xml, sizeof(xml)); + + std::array error; + error[0] = '\0'; + mjModel* model = mj_loadXML("model.xml", &vfs, error.data(), error.size()); + ASSERT_THAT(model, NotNull()); + + // warning should be in the error buffer + EXPECT_THAT(error.data(), HasSubstr("not rigid")); + + mj_deleteModel(model); + mj_deleteVFS(&vfs); +} + +TEST_F(MujocoTest, CompileWarningChainedToHandler) { + static constexpr char xml[] = R"( + + + + + + + + + + + )"; + std::array error; + mjSpec* spec = mj_parseXMLString(xml, 0, error.data(), error.size()); + ASSERT_THAT(spec, NotNull()) << error.data(); + + // install a custom log handler that captures warnings + std::vector captured_warnings; + static thread_local std::vector* capture_ptr = nullptr; + capture_ptr = &captured_warnings; + + // install custom log handler (replaces global, so mock is bypassed) + + mjfLogHandler prev = mju_setLogHandler([](const mjLogMessage* msg) { + if (msg->level == mjLOG_WARNING && capture_ptr) { + capture_ptr->push_back(msg->subject); + } + }); + + mjModel* model = mj_compile(spec, 0); + ASSERT_THAT(model, NotNull()); + + // restore log handler + mju_setLogHandler(prev); + capture_ptr = nullptr; + + // chaining should have forwarded warnings to our handler + EXPECT_THAT(captured_warnings, testing::Contains(HasSubstr("not rigid"))); + + mj_deleteModel(model); + mj_deleteSpec(spec); +} + } // namespace } // namespace mujoco diff --git a/test/user/user_objects_test.cc b/test/user/user_objects_test.cc index 3a5f8cff..b084a2ea 100644 --- a/test/user/user_objects_test.cc +++ b/test/user/user_objects_test.cc @@ -1181,7 +1181,7 @@ TEST_F(MjCGeomTest, IgnoreBadGeomOutsideInertiagrouprange) { TEST_F(MjCGeomTest, NanSize) { // even if the caller ignores warnings, models shouldn't compile with NaN // geom sizes - mju_user_warning = nullptr; + mock_warning_handler.ExpectWarnings(); static constexpr char xml[] = R"( diff --git a/test/xml/xml_write_read_test.cc b/test/xml/xml_write_read_test.cc index 0bc10d27..0c250d20 100644 --- a/test/xml/xml_write_read_test.cc +++ b/test/xml/xml_write_read_test.cc @@ -84,6 +84,12 @@ class WriteReadCompareTest : public XMLWriterTest, TEST_P(WriteReadCompareTest, WriteReadCompare) { std::string xml = GetParam(); + // If this is the flex_line_obj model, expect the 'is not rigid' warning + if (absl::StrContains(xml, "flex_line_obj")) { + EXPECT_CALL(mock_warning_handler, Warn(testing::HasSubstr("is not rigid"))) + .WillRepeatedly(testing::Return()); + } + // full precision float printing FullFloatPrecision increase_precision; diff --git a/wasm/codegen/generated/bindings.cc b/wasm/codegen/generated/bindings.cc index 74455c6e..5232bc98 100644 --- a/wasm/codegen/generated/bindings.cc +++ b/wasm/codegen/generated/bindings.cc @@ -115,15 +115,20 @@ EMSCRIPTEN_DECLARE_VAL_TYPE(StringOrNull); mju_error("Invalid argument: %s is undefined", #val); \ } -void ThrowMujocoErrorToJS(const char* msg) { - // Get a handle to the JS global Error constructor function, create a new - // object instance and then throw the object as an exception using the - // val::throw_() helper function. - val(val::global("Error").new_(val("MuJoCo Error: " + std::string(msg)))) - .throw_(); +void ThrowMujocoErrorToJS(const mjLogMessage* msg) { + if (msg->level == mjLOG_ERROR) { + std::string message = msg->subject; + if (msg->func) { + message = std::string(msg->func) + ": " + msg->subject; + } + // Get a handle to the JS global Error constructor function, create a new + // object instance and then throw the object as an exception using the + // val::throw_() helper function. + val(val::global("Error").new_(val("MuJoCo Error: " + message))).throw_(); + } } __attribute__((constructor)) void InitMuJoCoErrorHandler() { - mju_user_error = ThrowMujocoErrorToJS; + mju_setLogHandler(ThrowMujocoErrorToJS); } // Generates a descriptive error message for when a key lookup fails. @@ -8730,20 +8735,44 @@ std::unique_ptr parseXMLString_wrapper(const std::string &xml) { std::unique_ptr mj_compile_wrapper_1(const MjSpec& spec) { mjSpec* spec_ptr = spec.get(); + + // suppress stderr playback: warnings are raised via console.warn() below + mjfLogHandler prev = _mjPRIVATE_setTlsLogHandler([](const mjLogMessage*) {}); mjModel* model = mj_compile(spec_ptr, nullptr); - if (!model || mjs_isWarning(spec_ptr)) { + _mjPRIVATE_setTlsLogHandler(prev); + if (!model) { mju_error("%s", mjs_getError(spec_ptr)); } + int num_warnings = mjs_numWarnings(spec_ptr); + if (num_warnings > 0) { + for (int i = 0; i < num_warnings; ++i) { + val::global("console").call( + "warn", + val("MuJoCo Warning: " + std::string(mjs_getWarning(spec_ptr, i)))); + } + } return std::unique_ptr(new MjModel(model)); } std::unique_ptr mj_compile_wrapper_2(const MjSpec& spec, const MjVFS& vfs) { mjSpec* spec_ptr = spec.get(); mjVFS* vfs_ptr = vfs.get(); + + // suppress stderr playback: warnings are raised via console.warn() below + mjfLogHandler prev = _mjPRIVATE_setTlsLogHandler([](const mjLogMessage*) {}); mjModel* model = mj_compile(spec_ptr, vfs_ptr); - if (!model || mjs_isWarning(spec_ptr)) { + _mjPRIVATE_setTlsLogHandler(prev); + if (!model) { mju_error("%s", mjs_getError(spec_ptr)); } + int num_warnings = mjs_numWarnings(spec_ptr); + if (num_warnings > 0) { + for (int i = 0; i < num_warnings; ++i) { + val::global("console").call( + "warn", + val("MuJoCo Warning: " + std::string(mjs_getWarning(spec_ptr, i)))); + } + } return std::unique_ptr(new MjModel(model)); } @@ -10101,6 +10130,10 @@ std::optional mjs_getSpecDefault_wrapper(const MjSpec& s) { return MjsDefault(result); } +std::string mjs_getWarning_wrapper(const MjSpec& spec, int index) { + return std::string(mjs_getWarning(spec.get(), index)); +} + std::optional mjs_getWrap_wrapper(const MjsTendon& tendonspec, int i) { mjsWrap* result = mjs_getWrap(tendonspec.get(), i); if (result == nullptr) { @@ -10178,6 +10211,10 @@ std::optional mjs_nextElement_wrapper(const MjSpec& s, const MjsElem return MjsElement(result); } +int mjs_numWarnings_wrapper(const MjSpec& spec) { + return mjs_numWarnings(spec.get()); +} + std::string mjs_resolveOrientation_wrapper(const val& quat, mjtByte degree, const String& sequence, const MjsOrientation& orientation) { CHECK_VAL(sequence); UNPACK_VALUE(double, quat); @@ -13742,6 +13779,7 @@ EMSCRIPTEN_BINDINGS(mujoco_bindings) { function("mjs_getParent", &mjs_getParent_wrapper); function("mjs_getSpec", &mjs_getSpec_wrapper); function("mjs_getSpecDefault", &mjs_getSpecDefault_wrapper); + function("mjs_getWarning", &mjs_getWarning_wrapper); function("mjs_getWrap", &mjs_getWrap_wrapper); function("mjs_getWrapCoef", &mjs_getWrapCoef_wrapper); function("mjs_getWrapDivisor", &mjs_getWrapDivisor_wrapper); @@ -13753,6 +13791,7 @@ EMSCRIPTEN_BINDINGS(mujoco_bindings) { function("mjs_makeMesh", &mjs_makeMesh_wrapper); function("mjs_nextChild", &mjs_nextChild_wrapper); function("mjs_nextElement", &mjs_nextElement_wrapper); + function("mjs_numWarnings", &mjs_numWarnings_wrapper); function("mjs_resolveOrientation", &mjs_resolveOrientation_wrapper); function("mjs_sensorDim", &mjs_sensorDim_wrapper); function("mjs_setDeepCopy", &mjs_setDeepCopy_wrapper); diff --git a/wasm/codegen/templates/bindings.cc b/wasm/codegen/templates/bindings.cc index df506497..c9c61935 100644 --- a/wasm/codegen/templates/bindings.cc +++ b/wasm/codegen/templates/bindings.cc @@ -113,15 +113,20 @@ EMSCRIPTEN_DECLARE_VAL_TYPE(StringOrNull); mju_error("Invalid argument: %s is undefined", #val); \ } -void ThrowMujocoErrorToJS(const char* msg) { - // Get a handle to the JS global Error constructor function, create a new - // object instance and then throw the object as an exception using the - // val::throw_() helper function. - val(val::global("Error").new_(val("MuJoCo Error: " + std::string(msg)))) - .throw_(); +void ThrowMujocoErrorToJS(const mjLogMessage* msg) { + if (msg->level == mjLOG_ERROR) { + std::string message = msg->subject; + if (msg->func) { + message = std::string(msg->func) + ": " + msg->subject; + } + // Get a handle to the JS global Error constructor function, create a new + // object instance and then throw the object as an exception using the + // val::throw_() helper function. + val(val::global("Error").new_(val("MuJoCo Error: " + message))).throw_(); + } } __attribute__((constructor)) void InitMuJoCoErrorHandler() { - mju_user_error = ThrowMujocoErrorToJS; + mju_setLogHandler(ThrowMujocoErrorToJS); } // Generates a descriptive error message for when a key lookup fails. @@ -798,20 +803,44 @@ std::unique_ptr parseXMLString_wrapper(const std::string &xml) { std::unique_ptr mj_compile_wrapper_1(const MjSpec& spec) { mjSpec* spec_ptr = spec.get(); + + // suppress stderr playback: warnings are raised via console.warn() below + mjfLogHandler prev = _mjPRIVATE_setTlsLogHandler([](const mjLogMessage*) {}); mjModel* model = mj_compile(spec_ptr, nullptr); - if (!model || mjs_isWarning(spec_ptr)) { + _mjPRIVATE_setTlsLogHandler(prev); + if (!model) { mju_error("%s", mjs_getError(spec_ptr)); } + int num_warnings = mjs_numWarnings(spec_ptr); + if (num_warnings > 0) { + for (int i = 0; i < num_warnings; ++i) { + val::global("console").call( + "warn", + val("MuJoCo Warning: " + std::string(mjs_getWarning(spec_ptr, i)))); + } + } return std::unique_ptr(new MjModel(model)); } std::unique_ptr mj_compile_wrapper_2(const MjSpec& spec, const MjVFS& vfs) { mjSpec* spec_ptr = spec.get(); mjVFS* vfs_ptr = vfs.get(); + + // suppress stderr playback: warnings are raised via console.warn() below + mjfLogHandler prev = _mjPRIVATE_setTlsLogHandler([](const mjLogMessage*) {}); mjModel* model = mj_compile(spec_ptr, vfs_ptr); - if (!model || mjs_isWarning(spec_ptr)) { + _mjPRIVATE_setTlsLogHandler(prev); + if (!model) { mju_error("%s", mjs_getError(spec_ptr)); } + int num_warnings = mjs_numWarnings(spec_ptr); + if (num_warnings > 0) { + for (int i = 0; i < num_warnings; ++i) { + val::global("console").call( + "warn", + val("MuJoCo Warning: " + std::string(mjs_getWarning(spec_ptr, i)))); + } + } return std::unique_ptr(new MjModel(model)); }