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:
committed by
Copybara-Service
parent
55c6332f20
commit
6f8bb5ef55
@@ -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"(
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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
@@ -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
@@ -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;
|
||||
};
|
||||
|
||||
@@ -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.");
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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>
|
||||
|
||||
@@ -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;
|
||||
|
||||
|
||||
Reference in New Issue
Block a user