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) {