From 8679c9fc59feb566db7dd2eb054df8e607157cac Mon Sep 17 00:00:00 2001 From: Yuval Tassa Date: Mon, 26 Feb 2024 06:38:03 -0800 Subject: [PATCH] Fix null pointer deref for unknown class names. PiperOrigin-RevId: 610391077 Change-Id: I5abd62ff46e31e788f06fe111adb3c25287d4fdc --- src/user/user_api.cc | 12 ++++++++++++ src/user/user_api.h | 3 +++ src/xml/xml_native_reader.cc | 10 ++++++---- test/xml/xml_native_reader_test.cc | 26 +++++++++++++++++++++++++- 4 files changed, 46 insertions(+), 5 deletions(-) diff --git a/src/user/user_api.cc b/src/user/user_api.cc index cb735dcb..171a300f 100644 --- a/src/user/user_api.cc +++ b/src/user/user_api.cc @@ -346,6 +346,18 @@ mjmDefault* mjm_getDefault(mjElement element) { +// find default in model by class name +mjmDefault* mjm_findDefault(mjmModel* modelspec, const char* classname) { + mjCModel* model = reinterpret_cast(modelspec->element); + mjCDef* cdef = model->FindDef(classname); + if (!cdef) { + return nullptr; + } + return &cdef->spec; +} + + + // find body in model by name mjmBody* mjm_findBody(mjmModel* modelspec, const char* name) { mjCModel* model = reinterpret_cast(modelspec->element); diff --git a/src/user/user_api.h b/src/user/user_api.h index 5b86ccd5..971fb7a4 100644 --- a/src/user/user_api.h +++ b/src/user/user_api.h @@ -821,6 +821,9 @@ MJAPI mjmModel* mjm_getModel(mjmBody* body); // Get default corresponding to an mjElement. MJAPI mjmDefault* mjm_getDefault(mjElement element); +// Find default in model by class name. +MJAPI mjmDefault* mjm_findDefault(mjmModel* model, const char* classname); + // Find body in model by name. MJAPI mjmBody* mjm_findBody(mjmModel* model, const char* name); diff --git a/src/xml/xml_native_reader.cc b/src/xml/xml_native_reader.cc index 7e3fab47..b5ff9417 100644 --- a/src/xml/xml_native_reader.cc +++ b/src/xml/xml_native_reader.cc @@ -31,8 +31,8 @@ #include #include -#include #include +#include #include "engine/engine_plugin.h" #include "engine/engine_util_errmem.h" #include "engine/engine_util_misc.h" @@ -3986,12 +3986,14 @@ void mjXReader::Keyframe(XMLElement* section) { // get defaults class mjmDefault* mjXReader::GetClass(XMLElement* section) { string text; - mjmDefault* def = 0; + mjmDefault* def = nullptr; if (ReadAttrTxt(section, "class", text)) { - def = &model->FindDef(text)->spec; + def = mjm_findDefault(&model->spec, text.c_str()); if (!def) { - throw mjXError(section, "unknown default class"); + throw mjXError( + section, + std::string("unknown default class name '" + text + "'").c_str()); } } diff --git a/test/xml/xml_native_reader_test.cc b/test/xml/xml_native_reader_test.cc index ae45a4a9..1b37dd4c 100644 --- a/test/xml/xml_native_reader_test.cc +++ b/test/xml/xml_native_reader_test.cc @@ -33,12 +33,13 @@ namespace mujoco { namespace { using ::std::string; +using ::testing::AllOf; using ::testing::Eq; +using ::testing::FloatEq; using ::testing::HasSubstr; using ::testing::IsNan; using ::testing::IsNull; using ::testing::NotNull; -using ::testing::FloatEq; using XMLReaderTest = MujocoTest; @@ -462,6 +463,29 @@ TEST_F(XMLReaderTest, RepeatedDefaultName) { EXPECT_THAT(error.data(), HasSubstr("repeated default class name")); } +TEST_F(XMLReaderTest, InvalidDefaultClassName) { + static constexpr char xml[] = R"( + + + + + + + + + + + + + )"; + std::array error; + mjModel* model = LoadModelFromString(xml, error.data(), error.size()); + ASSERT_THAT(model, IsNull()) << error.data(); + EXPECT_THAT(error.data(), + AllOf(HasSubstr("unknown default class name 'invalid'"), + HasSubstr("Element 'geom'"), HasSubstr("line 10"))); +} + // ------------------------ test including ------------------------------------- // credit: https://www.mjt.me.uk/posts/smallest-png/