Add mj_markStack and mj_freeStack as public API functions.

Also add asan instrumentation to detect stack frame leakages (i.e. `mj_markStack` without a corresponding `mj_freeStack` in the same caller function).

PiperOrigin-RevId: 562625645
Change-Id: I4e3ff66ca0b9d08ed0a95cef45393db8e3053e22
This commit is contained in:
Saran Tunyasuvunakool
2023-09-04 17:03:23 -07:00
committed by Copybara-Service
parent 9308e1d383
commit 94a8705ad0
17 changed files with 311 additions and 45 deletions
+81 -3
View File
@@ -17,24 +17,26 @@
#include "src/engine/engine_io.h"
#include <array>
#include <climits>
#include <cstdint>
#include <cstdio>
#include <cstring>
#include <filesystem>
#include <fstream>
#include <iostream>
#include <string>
#include <vector>
#include <gmock/gmock.h>
#include <gtest/gtest.h>
#include <gtest/gtest-spi.h> // IWYU pragma: keep
#include <absl/strings/str_format.h>
#include <mujoco/mjxmacro.h>
#include <mujoco/mujoco.h>
#include "src/engine/engine_util_errmem.h"
#include "test/fixture.h"
namespace mujoco {
namespace {
using ::testing::ContainsRegex; // NOLINT(misc-unused-using-decls) asan only
using ::testing::HasSubstr;
using ::testing::IsNull;
using ::testing::NotNull;
@@ -737,5 +739,81 @@ TEST_F(ValidateReferencesTest, Tuples) {
mj_deleteModel(model);
}
TEST_F(EngineIoTest, CanMarkAndFreeStack) {
constexpr char xml[] = R"(
<mujoco>
<worldbody>
</worldbody>
</mujoco>
)";
std::array<char, 1024> error;
mjModel* model = LoadModelFromString(xml, error.data(), error.size());
ASSERT_THAT(model, NotNull()) << "Failed to load model: " << error.data();
mjData* data = mj_makeData(model);
ASSERT_THAT(data, NotNull());
auto pstack_before = data->pstack;
mj_markStack(data);
EXPECT_GT(data->pstack, pstack_before);
mj_freeStack(data);
EXPECT_EQ(data->pstack, pstack_before);
mj_deleteData(data);
mj_deleteModel(model);
}
#ifdef ADDRESS_SANITIZER
void MarkFreeStack(mjData* d, bool free) {
mj_markStack(d);
if (free) {
mj_freeStack(d);
}
}
TEST_F(EngineIoTest, CanDetectStackFrameLeakage) {
static constexpr char xml[] = R"(
<mujoco>
<worldbody>
</worldbody>
</mujoco>
)";
std::array<char, 1024> error;
mjModel* model = LoadModelFromString(xml, error.data(), error.size());
ASSERT_THAT(model, NotNull()) << "Failed to load model: " << error.data();
mjData* data = mj_makeData(model);
ASSERT_THAT(data, NotNull());
// MarkFreeStack correctly calls mj_freeStack, should not error.
MarkFreeStack(data, /* free= */ true);
// MarkFreeStack calls mj_markStack without mj_freeStack, the next call to
// mj_freeStack should detect the stack frame leakage.
mj_markStack(data);
MarkFreeStack(data, /* free= */ false);
EXPECT_THAT(
MjuErrorMessageFrom(mj_freeStack)(data),
ContainsRegex(
"mj_markStack in MarkFreeStack at .*engine_io_test\\.cc.* has no "
"corresponding mj_freeStack"));
// Dangling stack frames should be detected in mj_deleteData.
mj_resetData(model, data);
mj_markStack(data);
EXPECT_THAT(
MjuErrorMessageFrom(mj_deleteData)(data),
ContainsRegex(
"mj_markStack in .+EngineIoTest_CanDetectStackFrameLeakage.+ has no "
"corresponding mj_freeStack"));
mj_resetData(model, data);
mj_deleteData(data);
mj_deleteModel(model);
}
#endif
} // namespace
} // namespace mujoco
+12 -6
View File
@@ -27,18 +27,21 @@ namespace {
using testing::NotNull;
template <typename T, int N, int Capacity>
template <int prev_size, typename T, int N, int capacity>
constexpr int GetExpectedStackUsageBytes() {
if constexpr (N <= 0) {
return 0;
return prev_size;
} else {
constexpr auto RoundUpToAlignment =
[](int x, int alignment) {
return alignment * (x / alignment + ((x % alignment) ? 1 : 0));
};
return RoundUpToAlignment(sizeof(mjArrayList), alignof(mjArrayList)) +
RoundUpToAlignment(Capacity * sizeof(T), alignof(std::max_align_t)) +
GetExpectedStackUsageBytes<T, N - Capacity, 2 * Capacity>();
constexpr int size_with_arraylist = RoundUpToAlignment(
prev_size + sizeof(mjArrayList), alignof(mjArrayList));
constexpr int size_with_buffer = RoundUpToAlignment(
size_with_arraylist + capacity * sizeof(T), alignof(std::max_align_t));
return GetExpectedStackUsageBytes<size_with_buffer, T, N - capacity,
2 * capacity>();
}
}
@@ -48,6 +51,7 @@ TEST(TestMjArrayList, TestMjArrayListSingleThreaded) {
ASSERT_THAT(m, NotNull()) << "Failed to load model: " << error.data();
mjData* d = mj_makeData(m);
mjMARKSTACK;
using DataType = int;
constexpr int kInitialCapacity = 10;
mjArrayList* array_list =
@@ -59,8 +63,10 @@ TEST(TestMjArrayList, TestMjArrayListSingleThreaded) {
}
EXPECT_EQ(mju_arrayListSize(array_list), kNumElements);
constexpr int kFrameMarkerSize = 2 * sizeof(size_t) + sizeof(void*);
constexpr int kExpectedMaxUseStack =
GetExpectedStackUsageBytes<DataType, kNumElements, kInitialCapacity>();
GetExpectedStackUsageBytes<kFrameMarkerSize, DataType, kNumElements,
kInitialCapacity>();
EXPECT_EQ(d->maxuse_stack, kExpectedMaxUseStack);
for (int i = 0; i < kNumElements; ++i) {
+35
View File
@@ -15,12 +15,23 @@
#ifndef MUJOCO_TEST_FIXTURE_H_
#define MUJOCO_TEST_FIXTURE_H_
#include <csetjmp>
#include <cstring>
#include <optional>
#include <string>
#include <vector>
#include <gtest/gtest.h>
#include <absl/strings/string_view.h>
#include <mujoco/mjdata.h>
#include <mujoco/mjmodel.h>
#include <mujoco/mujoco.h>
extern "C" {
MJAPI void _mjPRIVATE__set_tls_error_fn(decltype(mju_user_error));
MJAPI decltype(mju_user_error) _mjPRIVATE__get_tls_error_fn();
}
namespace mujoco {
// Installs and uninstalls error callbacks on MuJoCo that fail the currently
@@ -44,6 +55,30 @@ class MujocoTest : public ::testing::Test {
MujocoErrorTestGuard error_guard;
};
template <typename Return, typename... Args>
auto MjuErrorMessageFrom(Return (*func)(Args...)) {
thread_local std::jmp_buf current_jmp_buf;
thread_local char err_msg[1000];
auto* old_error_handler = _mjPRIVATE__get_tls_error_fn();
auto* new_error_handler = +[](const char* msg) -> void {
std::strncpy(err_msg, msg, sizeof(err_msg));
std::longjmp(current_jmp_buf, 1);
};
return [func, old_error_handler,
new_error_handler](Args... args) -> std::string {
if (setjmp(current_jmp_buf) == 0) {
err_msg[0] = '\0';
_mjPRIVATE__set_tls_error_fn(new_error_handler);
func(args...);
}
_mjPRIVATE__set_tls_error_fn(old_error_handler);
return err_msg;
};
}
// Returns a path to a data file, under the mujoco/test directory.
const std::string GetTestDataFilePath(absl::string_view path);
+3 -2
View File
@@ -17,6 +17,7 @@
#include <array>
#include <limits>
#include <string>
#include <vector>
#include <gmock/gmock.h>
#include <gtest/gtest.h>
@@ -43,12 +44,12 @@ TEST_F(XMLReaderTest, MemorySize) {
{
static constexpr char xml[] = R"(
<mujoco>
<size memory="256"/>
<size memory="512"/>
</mujoco>
)";
mjModel* model = LoadModelFromString(xml, error.data(), error.size());
ASSERT_THAT(model, NotNull()) << error.data();
EXPECT_EQ(model->narena, 256);
EXPECT_EQ(model->narena, 512);
mj_deleteModel(model);
}
{