Fix out-of-bound read in mju_combineSparseCount.

PiperOrigin-RevId: 535615505
Change-Id: I740d4f50b46ec2a1cd1a5a615fb4312f1bbebd94
This commit is contained in:
Saran Tunyasuvunakool
2023-05-26 07:43:59 -07:00
committed by Copybara-Service
parent cce7fb04d1
commit 1d79657512
3 changed files with 97 additions and 26 deletions
+17 -26
View File
@@ -1255,35 +1255,26 @@ void mj_makeImpedance(const mjModel* m, mjData* d) {
//------------------------------------- constraint counting ----------------------------------------
// count the number of non-zeros in the sum of two sparse vectors
static int mju_combineSparseCount(int a_nnz, int b_nnz, const int* a_ind, const int* b_ind) {
int c_nnz, d_nnz;
const int* c_ind;
const int* d_ind;
int mju_combineSparseCount(int a_nnz, int b_nnz, const int* a_ind, const int* b_ind) {
int a = 0;
int b = 0;
int nnz = 0;
// choose c to have the least number of non-zeros
if (b_nnz<a_nnz) {
c_nnz = b_nnz;
c_ind = b_ind;
d_nnz = a_nnz;
d_ind = a_ind;
} else {
c_nnz = a_nnz;
c_ind = a_ind;
d_nnz = b_nnz;
d_ind = b_ind;
}
int nnz=d_nnz, j=0;
for (int i=0; i<c_nnz; i++) {
while (d_ind[j]<c_ind[i]) {
j++;
}
if (d_ind[j]>c_ind[i]) {
nnz++;
}
// while there are elements remaining in both a_ind and b_ind
while (a < a_nnz && b < b_nnz) {
// add the smaller element of either a_ind[a] or b_ind[b] to the combined nnz
++nnz;
// if a_ind[a] == b_ind[b], increment both a and b so that we don't double count
// otherwise, increment the index pointing to the smaller element
int aa = a;
int bb = b;
if (a_ind[aa] <= b_ind[bb]) ++a;
if (a_ind[aa] >= b_ind[bb]) ++b;
}
// count remaining elements from the vector with larger nnz
nnz += (a_nnz - a) + (b_nnz - b);
return nnz;
}
+6
View File
@@ -93,6 +93,12 @@ void mj_diagApprox(const mjModel* m, mjData* d);
void mj_makeImpedance(const mjModel* m, mjData* d);
//------------------------- constraint counting
// count the number of non-zeros in the sum of two sparse vectors
MJAPI int mju_combineSparseCount(int a_nnz, int b_nnz, const int* a_ind, const int* b_ind);
//---------------------------- top-level API for constraint construction ---------------------------
// main driver: call all functions above
@@ -14,6 +14,7 @@
// Tests for engine/engine_core_constraint.c.
#include <array>
#include <cstddef>
#include <string>
@@ -21,6 +22,7 @@
#include <gtest/gtest.h>
#include <mujoco/mjmodel.h>
#include <mujoco/mujoco.h>
#include "src/engine/engine_core_constraint.h"
#include "src/engine/engine_support.h"
#include "test/fixture.h"
@@ -195,5 +197,77 @@ TEST_F(CoreConstraintTest, JacobianPreAllocate) {
}
}
TEST_F(CoreConstraintTest, CombineSparseCount) {
{
std::array a_ind{0, 1};
std::array b_ind{2};
EXPECT_EQ(mju_combineSparseCount(
a_ind.size(), b_ind.size(), a_ind.data(), b_ind.data()), 3);
}
{
std::array a_ind{2};
std::array b_ind{0, 1};
EXPECT_EQ(mju_combineSparseCount(
a_ind.size(), b_ind.size(), a_ind.data(), b_ind.data()), 3);
}
{
std::array a_ind{0, 1};
std::array b_ind{2, 3, 4};
EXPECT_EQ(mju_combineSparseCount(
a_ind.size(), b_ind.size(), a_ind.data(), b_ind.data()), 5);
}
{
std::array a_ind{5, 6};
std::array b_ind{1, 3, 8};
EXPECT_EQ(mju_combineSparseCount(
a_ind.size(), b_ind.size(), a_ind.data(), b_ind.data()), 5);
}
{
std::array a_ind{1, 2, 3};
std::array b_ind{0, 4};
EXPECT_EQ(mju_combineSparseCount(
a_ind.size(), b_ind.size(), a_ind.data(), b_ind.data()), 5);
}
{
std::array a_ind{1, 4};
std::array b_ind{2, 3};
EXPECT_EQ(mju_combineSparseCount(
a_ind.size(), b_ind.size(), a_ind.data(), b_ind.data()), 4);
}
{
std::array a_ind{0, 1, 3};
std::array b_ind{0, 3, 4};
EXPECT_EQ(mju_combineSparseCount(
a_ind.size(), b_ind.size(), a_ind.data(), b_ind.data()), 4);
}
{
std::array a_ind{1, 3, 5, 6};
std::array b_ind{1, 3, 5, 6};
EXPECT_EQ(mju_combineSparseCount(
a_ind.size(), b_ind.size(), a_ind.data(), b_ind.data()), 4);
}
EXPECT_EQ(mju_combineSparseCount(0, 0, nullptr, nullptr), 0);
{
std::array b_ind{1, 2};
EXPECT_EQ(
mju_combineSparseCount(0, b_ind.size(), nullptr, b_ind.data()), 2);
}
{
std::array a_ind{0};
EXPECT_EQ(
mju_combineSparseCount(a_ind.size(), 0, a_ind.data(), nullptr), 1);
}
}
} // namespace
} // namespace mujoco