Check that alignment is power of 2 in mj_stackAlloc.

PiperOrigin-RevId: 562024496
Change-Id: I9230c6dc05ebb78e935f9b2f606aa77c921d67e5
This commit is contained in:
Saran Tunyasuvunakool
2023-09-01 12:44:53 -07:00
committed by Copybara-Service
parent 29aa5e4a41
commit 33f51856c1
7 changed files with 34 additions and 24 deletions
+1 -1
View File
@@ -1239,7 +1239,7 @@ mj_stackAlloc
.. mujoco-include:: mj_stackAlloc
Allocate a number of bytes on :ref:`mjData` stack at a specific alignment which must be a power of 2.
Allocate a number of bytes on :ref:`mjData` stack at a specific alignment.
Call mju_error on stack overflow.
.. _mj_stackAllocNum:
+1 -1
View File
@@ -188,7 +188,7 @@ MJAPI void mj_resetDataDebug(const mjModel* m, mjData* d, unsigned char debug_va
// Reset data, set fields from specified keyframe.
MJAPI void mj_resetDataKeyframe(const mjModel* m, mjData* d, int key);
// Allocate a number of bytes on mjData stack at a specific alignment which must be a power of 2.
// Allocate a number of bytes on mjData stack at a specific alignment.
// Call mju_error on stack overflow.
MJAPI void* mj_stackAlloc(mjData* d, size_t bytes, size_t alignment);
+1 -1
View File
@@ -699,7 +699,7 @@ FUNCTIONS: Mapping[str, FunctionDecl] = dict([
type=ValueType(name='size_t'),
),
),
doc='Allocate a number of bytes on mjData stack at a specific alignment which must be a power of 2. Call mju_error on stack overflow.', # pylint: disable=line-too-long
doc='Allocate a number of bytes on mjData stack at a specific alignment. Call mju_error on stack overflow.', # pylint: disable=line-too-long
)),
('mj_stackAllocNum',
FunctionDecl(
+1 -1
View File
@@ -948,7 +948,7 @@ Euler integrator, semi-implicit in velocity.
def test_can_raise_error(self):
self.data.pstack = self.data.narena
with self.assertRaisesRegex(mujoco.FatalError,
r'\AmjData stack overflow'):
r'\Amj_stackAlloc: insufficient memory:'):
mujoco.mj_forward(self.model, self.data)
def test_mjcb_time(self):
+14 -8
View File
@@ -25,7 +25,7 @@
#endif
// IWYU pragma: end_keep
// Windows
// Sorting and case-insensitive comparison functions.
#ifdef _WIN32
#define strcasecmp _stricmp
#define strncasecmp _strnicmp
@@ -34,20 +34,15 @@
qsort_s(buf, elnum, elsz, func, context)
#define quicksortfunc(name, context, el1, el2) \
static int name(void* context, const void* el1, const void* el2)
// Unix-common
#else
#else // assumes POSIX
#include <strings.h>
// Apple
#ifdef __APPLE__
#define mjQUICKSORT(buf, elnum, elsz, func, context) \
qsort_r(buf, elnum, elsz, context, func)
#define quicksortfunc(name, context, el1, el2) \
static int name(void* context, const void* el1, const void* el2)
// non-Apple
#else
#else // non-Apple
#define mjQUICKSORT(buf, elnum, elsz, func, context) \
qsort_r(buf, elnum, elsz, func, context)
#define quicksortfunc(name, context, el1, el2) \
@@ -55,6 +50,7 @@
#endif
#endif
// Switch-case fallthrough annotation.
#if defined(__cplusplus)
#define mjFALLTHROUGH [[fallthrough]]
#elif defined(__clang__) || (defined(__GNUC__) && __GNUC__ >= 7)
@@ -63,10 +59,20 @@
#define mjFALLTHROUGH ((void) 0)
#endif
// MSVC only provides max_align_t in C++.
#if defined(_MSC_VER) && !defined(__clang__) && !defined(__cplusplus)
typedef long double mjtMaxAlign;
#else
typedef max_align_t mjtMaxAlign;
#endif
// Branch prediction hints.
#if defined(__GNUC__)
#define mjLIKELY(x) __builtin_expect(!!(x), 1)
#define mjUNLIKELY(x) __builtin_expect(!!(x), 0)
#else
#define mjLIKELY(x) (x)
#define mjUNLIKELY(x) (x)
#endif
#endif // MUJOCO_SRC_ENGINE_ENGINE_CROSSPLATFORM_H_
+14 -10
View File
@@ -26,7 +26,7 @@
#include <mujoco/mjmacro.h>
#include <mujoco/mjplugin.h>
#include <mujoco/mjxmacro.h>
#include "engine/engine_array_safety.h" // IWYU pragma: keep
#include "engine/engine_crossplatform.h"
#include "engine/engine_resource.h"
#include "engine/engine_macro.h"
#include "engine/engine_plugin.h"
@@ -45,9 +45,13 @@
static const int MAX_ARRAY_SIZE = INT_MAX / 4;
// compute a % b assuming that the second argument is a power of 2
static inline size_t modpow2(size_t a, size_t b) {
return a & (b - 1);
// compute a % b with a fast code path if the second argument is a power of 2
static inline size_t fastmod(size_t a, size_t b) {
// (b & (b - 1)) == 0 implies that b is a power of 2
if (mjLIKELY((b & (b - 1)) == 0)) {
return a & (b - 1);
}
return a % b;
}
//------------------------------ mjLROpt -----------------------------------------------------------
@@ -1187,12 +1191,12 @@ mjData* mj_copyData(mjData* dest, const mjModel* m, const mjData* src) {
// allocate memory from the mjData arena
void* mj_arenaAlloc(mjData* d, size_t bytes, size_t alignment) {
size_t misalignment = modpow2(d->parena, alignment);
size_t misalignment = fastmod(d->parena, alignment);
size_t padding = misalignment ? alignment - misalignment : 0;
// check size
size_t bytes_available = d->narena - d->pstack;
if (d->parena + padding + bytes > bytes_available) {
if (mjUNLIKELY(d->parena + padding + bytes > bytes_available)) {
return NULL;
}
@@ -1218,7 +1222,7 @@ void* mj_arenaAlloc(mjData* d, size_t bytes, size_t alignment) {
// declared inline so that modular arithmetic with specific alignments can be optimized out
static inline void* stackalloc(mjData* d, size_t size, size_t alignment) {
// return NULL if empty
if (!size) {
if (mjUNLIKELY(!size)) {
return NULL;
}
@@ -1242,7 +1246,7 @@ static inline void* stackalloc(mjData* d, size_t size, size_t alignment) {
uintptr_t start_ptr = end_ptr - (size + mjREDZONE);
// align the pointer
start_ptr -= modpow2(start_ptr, alignment);
start_ptr -= fastmod(start_ptr, alignment);
// new top of the stack
uintptr_t new_pstack_ptr = start_ptr - mjREDZONE;
@@ -1255,8 +1259,8 @@ static inline void* stackalloc(mjData* d, size_t size, size_t alignment) {
// check size
size_t stack_available_bytes = end_ptr - ((uintptr_t)d->arena + d->parena);
size_t stack_required_bytes = end_ptr - new_pstack_ptr;
if (stack_required_bytes > stack_available_bytes) {
mju_error("mjData stack overflow: max = %zu, available = %zu, requested = %zu "
if (mjUNLIKELY(stack_required_bytes > stack_available_bytes)) {
mju_error("mj_stackAlloc: insufficient memory: max = %zu, available = %zu, requested = %zu "
"(ne = %d, nf = %d, nefc = %d, ncon = %d)",
stack_size_bytes, stack_available_bytes, stack_required_bytes,
d->ne, d->nf, d->nefc, d->ncon);
+2 -2
View File
@@ -102,10 +102,10 @@ MJAPI void mj_resetDataDebug(const mjModel* m, mjData* d, unsigned char debug_va
// reset data, set fields from specified keyframe
MJAPI void mj_resetDataKeyframe(const mjModel* m, mjData* d, int key);
// mjData arena allocate (alignment must be a power of 2)
// mjData arena allocate
MJAPI void* mj_arenaAlloc(mjData* d, size_t bytes, size_t alignment);
// mjData stack allocate (alignment must be a power of 2)
// mjData stack allocate
MJAPI void* mj_stackAlloc(mjData* d, size_t bytes, size_t alignment);
// mjData stack allocate for array of mjtNums