Refactor compiler warning handling.

Compiler warnings are now accumulated in a vector of strings within the mjSpec object. New API functions `mjs_numWarnings` and `mjs_getWarning` are added to access these warnings. The compiler's log handler now chains warnings to the global log handler, ensuring they are still displayed immediately. Call sites in `mj_loadXML`, `mj_compile`, and the Python and WASM bindings have been updated to use the new warning API.

PiperOrigin-RevId: 933361650
Change-Id: I47cab98a460c57b0898c0a1a43fce2a5b9648eb1
This commit is contained in:
Yuval Tassa
2026-06-16 16:27:52 -07:00
committed by Copybara-Service
parent 55c6332f20
commit 6f8bb5ef55
25 changed files with 640 additions and 89 deletions
+3 -1
View File
@@ -80,7 +80,9 @@ struct ActLimitedTestCase {
mjtIntegrator integrator;
};
using ParametrizedForwardTest = ::testing::TestWithParam<ActLimitedTestCase>;
class ParametrizedForwardTest
: public MujocoTest,
public ::testing::WithParamInterface<ActLimitedTestCase> {};
TEST_P(ParametrizedForwardTest, ActLimited) {
static constexpr char xml[] = R"(
+4
View File
@@ -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);
+4 -2
View File
@@ -45,7 +45,9 @@ constexpr int GetExpectedStackUsageBytes() {
}
}
TEST(TestMjArrayList, TestMjArrayListSingleThreaded) {
class TestMjArrayList : public MujocoTest {};
TEST_F(TestMjArrayList, TestMjArrayListSingleThreaded) {
std::array<char, 1024> error;
mjModel* m = LoadModelFromString("<mujoco/>", 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("<mujoco/>", error, sizeof(error));
ASSERT_THAT(m, NotNull()) << "Failed to load model: " << error;
+99 -13
View File
@@ -24,7 +24,6 @@
#include <sstream>
#include <string>
#include <string_view>
#include <utility>
#include <vector>
@@ -33,6 +32,7 @@
#include <absl/base/attributes.h>
#include <absl/base/const_init.h>
#include <absl/base/thread_annotations.h>
#include <absl/strings/match.h>
#include <absl/strings/str_cat.h>
#include <absl/strings/str_join.h>
#include <absl/synchronization/mutex.h>
@@ -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;
}
+26 -1
View File
@@ -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;
};
+7
View File
@@ -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.");
}
+163
View File
@@ -3291,5 +3291,168 @@ TEST_F(MujocoTest, CompilerTimers) {
mj_deleteSpec(spec);
}
// -------------------- test compile warning infrastructure --------------------
TEST_F(MujocoTest, CompileWarningCount) {
static constexpr char xml[] = R"(
<mujoco>
<worldbody>
<body name="parent">
<geom size="1"/>
<flexcomp name="grid" type="grid" count="3 3 1" spacing="0.1 0.1 0.1"
dim="2" radius="0.01">
<contact internal="false"/>
</flexcomp>
</body>
</worldbody>
</mujoco>
)";
std::array<char, 1024> 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"(
<mujoco>
<worldbody>
<body name="parent">
<geom size="1"/>
<flexcomp name="grid" type="grid" count="3 3 1" spacing="0.1 0.1 0.1"
dim="2" radius="0.01">
<contact internal="false"/>
</flexcomp>
</body>
</worldbody>
</mujoco>
)";
std::array<char, 1024> 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"(
<mujoco>
<worldbody>
<body name="parent">
<geom size="1"/>
<flexcomp name="grid" type="grid" count="3 3 1" spacing="0.1 0.1 0.1"
dim="2" radius="0.01">
<contact internal="false"/>
</flexcomp>
</body>
</worldbody>
</mujoco>
)";
// write xml to VFS
mjVFS vfs;
mj_defaultVFS(&vfs);
mj_addBufferVFS(&vfs, "model.xml", xml, sizeof(xml));
std::array<char, 1024> 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"(
<mujoco>
<worldbody>
<body name="parent">
<geom size="1"/>
<flexcomp name="grid" type="grid" count="3 3 1" spacing="0.1 0.1 0.1"
dim="2" radius="0.01">
<contact internal="false"/>
</flexcomp>
</body>
</worldbody>
</mujoco>
)";
std::array<char, 1024> 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<std::string> captured_warnings;
static thread_local std::vector<std::string>* 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
+1 -1
View File
@@ -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"(
<mujoco>
<worldbody>
+6
View File
@@ -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;