From 389736f5a185e2bfaf114aaa8d412f3b2ab500c1 Mon Sep 17 00:00:00 2001 From: Kyle Bayes Date: Tue, 29 Aug 2023 06:13:26 -0700 Subject: [PATCH] Add Base64, defined by RFC 4648, encoding and decoding to MuJoCo engine. PiperOrigin-RevId: 561008882 Change-Id: I494eb163bc82e8a5c9ad2116ef119072b08a9f48 --- src/engine/engine_util_misc.c | 149 ++++++++++++++++++- src/engine/engine_util_misc.h | 16 ++ test/engine/engine_util_misc_test.cc | 211 +++++++++++++++++++++++++++ 3 files changed, 375 insertions(+), 1 deletion(-) diff --git a/src/engine/engine_util_misc.c b/src/engine/engine_util_misc.c index 27c3be80..dd9a9cbe 100644 --- a/src/engine/engine_util_misc.c +++ b/src/engine/engine_util_misc.c @@ -14,7 +14,9 @@ #include "engine/engine_util_misc.h" +#include #include +#include #include #include #include @@ -595,6 +597,151 @@ mjtNum mju_muscleDynamics(mjtNum ctrl, mjtNum act, const mjtNum prm[3]) { +//---------------------------------------- Base64 -------------------------------------------------- + +// decoding function for Base64 +static uint32_t _decode(char ch) { + if (ch >= 'A' && ch <= 'Z') { + return ch - 'A'; + } + + if (ch >= 'a' && ch <= 'z') { + return (ch - 'a') + 26; + } + + if (ch >= '0' && ch <= '9') { + return (ch - '0') + 52; + } + + if (ch == '+') { + return 62; + } + + if (ch == '/') { + return 63; + } + + return 0; +} + + + +// encode data as Base64 into buf (including padding and null char) +// returns number of chars written in buf: 4 * [(ndata + 2) / 3] + 1 +size_t mju_encodeBase64(char* buf, const uint8_t* data, size_t ndata) { + static const char *table = + "ABCDEFGHIJKLMNOPQRSTUBWXYZabcdefghijklmnopqrstuvwxyz0123456789+/"; + + int i = 0, j = 0; + + // loop over 24 bit chunks + while (i + 3 <= ndata) { + // take next 24 bit chunk (3 bytes) + uint32_t byte_1 = data[i++]; + uint32_t byte_2 = data[i++]; + uint32_t byte_3 = data[i++]; + + // merge bytes into one 32 bit int + uint32_t k = (byte_1 << 16) | (byte_2 << 8) | byte_3; + + // encode 6 bit chucks into four chars + buf[j++] = table[(k >> 18) & 63]; + buf[j++] = table[(k >> 12) & 63]; + buf[j++] = table[(k >> 6) & 63]; + buf[j++] = table[(k >> 0) & 63]; + } + + // one byte left + if (i + 1 == ndata) { + uint32_t byte_1 = data[i]; + uint32_t k = byte_1 << 16; + buf[j++] = table[(k >> 18) & 63]; + buf[j++] = table[(k >> 12) & 63]; + buf[j++] = '='; // padding + buf[j++] = '='; // padding + } + + // two bytes left + if (i + 2 == ndata) { + uint32_t byte_1 = data[i++]; + uint32_t byte_2 = data[i]; + + uint32_t k = (byte_1 << 16) + (byte_2 << 8); + + buf[j++] = table[(k >> 18) & 63]; + buf[j++] = table[(k >> 12) & 63]; + buf[j++] = table[(k >> 6) & 63]; + buf[j++] = '='; // padding + } + + buf[j] = '\0'; + return 4 * ((ndata + 2) / 3) + 1; +} + + + +// return size in decoded bytes if s is a valid Base64 encoding +// return 0 if s is empty or invalid Base64 encoding +size_t mju_isValidBase64(const char* s) { + size_t i = 0; + int pad = 0; // 0, 1, or 2 zero padding at the end of s + + // validate chars + for (; s[i] && s[i] != '='; i++) { + if (!isalnum(s[i]) && s[i] != '/' && s[i] != '+') { + return 0; + } + } + + // padding at end + if (s[i] == '=') { + if (!s[i + 1]) { + pad = 1; // one '=' padding at end + } else if (s[i + 1] == '=' && !s[i + 2]) { + pad = 2; // two '=' padding at end + } else { + return 0; + } + } + + // strlen(s) must be a multiple of 4 + int len = i + pad; + return len % 4 ? 0 : 3 * (len / 4) - pad; +} + + + +// decode valid Base64 in string s into buf, undefined behavior if s is not valid Base64 +// returns number of bytes decoded (upper limit of 3 * (strlen(s) / 4)) +size_t mju_decodeBase64(uint8_t* buf, const char* s) { + size_t i = 0, j = 0; + + // loop over 24 bit chunks + while (s[i] != '\0') { + // take next 24 bit chuck (4 chars; 6 bits each) + uint32_t char_1 = _decode(s[i++]); + uint32_t char_2 = _decode(s[i++]); + uint32_t char_3 = _decode(s[i++]); + uint32_t char_4 = _decode(s[i++]); + + // merge into 32 bit int + uint32_t k = (char_1 << 18) | (char_2 << 12) | (char_3 << 6) | char_4; + + + // write up to three bytes (exclude padding at end) + buf[j++] = (k >> 16) & 0xFF; + if (s[i - 2] != '=') { + buf[j++] = (k >> 8) & 0xFF; + } + if (s[i - 1] != '=') { + buf[j++] = k & 0xFF; + } + } + return j; +} + + + //------------------------------ miscellaneous ----------------------------------------------------- // convert contact force to pyramid representation @@ -603,7 +750,7 @@ mjtNum mju_muscleDynamics(mjtNum ctrl, mjtNum act, const mjtNum prm[3]) { void mju_encodePyramid(mjtNum* pyramid, const mjtNum* force, const mjtNum* mu, int dim) { mjtNum a = force[0]/(dim-1), b; - // arbitary redundancy resolution: + // arbitrary redundancy resolution: // pyramid0_i + pyramid1_i = force_normal/(dim-1) = a // pyramid0_i - pyramid1_i = force_tangent_i/mu_i = b for (int i=0; i < dim-1; i++) { diff --git a/src/engine/engine_util_misc.h b/src/engine/engine_util_misc.h index 6ac874fc..6260d0c5 100644 --- a/src/engine/engine_util_misc.h +++ b/src/engine/engine_util_misc.h @@ -17,12 +17,14 @@ #include #include +#include #ifdef __cplusplus extern "C" { #endif #include +#include //------------------------------ tendons and actuators --------------------------------------------- @@ -49,6 +51,20 @@ MJAPI mjtNum mju_muscleDynamics(mjtNum ctrl, mjtNum act, const mjtNum prm[3]); // all 3 semi-axes of a geom MJAPI void mju_geomSemiAxes(const mjModel* m, int geom_id, mjtNum semiaxes[3]); +// ----------------------------- Base64 ----------------------------------------------------------- + +// encode data as Base64 into buf (including padding and null char) +// returns number of chars written in buf: 4 * [(ndata + 2) / 3] + 1 +MJAPI size_t mju_encodeBase64(char* buf, const uint8_t* data, size_t ndata); + +// return size in decoded bytes if s is a valid Base64 encoding +// return 0 if s is empty or invalid Base64 encoding +MJAPI size_t mju_isValidBase64(const char* s); + +// decode valid Base64 in string s into buf, undefined behavior if s is not valid Base64 +// returns number of bytes decoded (upper limit of 3 * (strlen(s) / 4)) +MJAPI size_t mju_decodeBase64(uint8_t* buf, const char* s); + //------------------------------ miscellaneous ---------------------------------------------------- // convert contact force to pyramid representation diff --git a/test/engine/engine_util_misc_test.cc b/test/engine/engine_util_misc_test.cc index 5181e67d..71702b8c 100644 --- a/test/engine/engine_util_misc_test.cc +++ b/test/engine/engine_util_misc_test.cc @@ -14,6 +14,11 @@ // Tests for engine/engine_util_solve.c. +#include +#include +#include +#include + #include #include #include @@ -28,6 +33,7 @@ using ::testing::DoubleNear; using ::testing::HasSubstr; using ::testing::Ne; using ::testing::StrEq; +using ::testing::ElementsAreArray; TEST_F(MujocoTest, PrintsMemoryWarning) { EXPECT_THAT(mju_warningText(mjWARN_CNSTRFULL, pow(2, 10)), @@ -215,5 +221,210 @@ TEST_F(MujocoTest, mju_makefullname_error5) { EXPECT_THAT(n, Ne(0)); } +// --------------------------------- Base64 ------------------------------------ + +using Base64Test = MujocoTest; + +TEST_F(Base64Test, mju_encodeBase64) { + std::array buffer; + std::array arr = {15, 134, 190, 255, 240}; + + std::size_t n = mju_encodeBase64(buffer.data(), arr.data(), arr.size()); + + EXPECT_THAT(buffer.data(), StrEq("D4a+//A=")); + EXPECT_THAT(n, std::strlen(buffer.data()) + 1); + EXPECT_THAT(n, buffer.size()); + +} + +TEST_F(Base64Test, mju_encodeBase64_align0) { + std::array buffer; + std::array arr = {'A', 'B', 'C'}; + + std::size_t n = mju_encodeBase64(buffer.data(), arr.data(), arr.size()); + + EXPECT_THAT(buffer.data(), StrEq("QUJD")); + EXPECT_THAT(n, std::strlen(buffer.data()) + 1); + EXPECT_THAT(n, buffer.size()); +} + +TEST_F(Base64Test, mju_encodeBase64_align1) { + std::array buffer; + std::array arr = {'A', 'B'}; + + std::size_t n = mju_encodeBase64(buffer.data(), arr.data(), arr.size()); + + EXPECT_THAT(buffer.data(), StrEq("QUI=")); + EXPECT_THAT(n, std::strlen(buffer.data()) + 1); + EXPECT_THAT(n, buffer.size()); + +} + +TEST_F(Base64Test, mju_encodeBase64_align2) { + std::array buffer; + std::array arr = {'A'}; + + std::size_t n = mju_encodeBase64(buffer.data(), arr.data(), arr.size()); + + EXPECT_THAT(buffer.data(), StrEq("QQ==")); + EXPECT_THAT(n, std::strlen(buffer.data()) + 1); + EXPECT_THAT(n, buffer.size()); +} + +TEST_F(Base64Test, mju_encodeBase64_null) { + std::array buffer; + + std::size_t n = mju_encodeBase64(buffer.data(), NULL, 0); + + EXPECT_THAT(n, 1); + EXPECT_THAT(buffer[0], '\0'); +} + +TEST_F(Base64Test, mju_encodeBase64_ones) { + std::array buffer; + std::array arr = {255, 255, 255}; + + std::size_t n = mju_encodeBase64(buffer.data(), arr.data(), arr.size()); + + EXPECT_THAT(buffer.data(), StrEq("////")); + EXPECT_THAT(n, std::strlen(buffer.data()) + 1); + EXPECT_THAT(n, buffer.size()); +} + +TEST_F(Base64Test, mju_isValidBase64_emptyStr) { + std::size_t n = mju_isValidBase64(""); + + EXPECT_THAT(n, 0); +} + +TEST_F(Base64Test, mju_isValidBase64_invalid1) { + std::size_t n = mju_isValidBase64("A"); + + EXPECT_THAT(n, 0); +} + +TEST_F(Base64Test, mju_isValidBase64_invalid2) { + std::size_t n = mju_isValidBase64("AAA"); + + EXPECT_THAT(n, 0); +} + +TEST_F(Base64Test, mju_isValidBase64_invalid3) { + std::size_t n = mju_isValidBase64("A==A"); + + EXPECT_THAT(n, 0); +} + +TEST_F(Base64Test, mju_isValidBase64_invalid5) { + std::size_t n = mju_isValidBase64("A==="); + + EXPECT_THAT(n, 0); +} + +TEST_F(Base64Test, mju_isValidBase64_invalid6) { + std::size_t n = mju_isValidBase64("aaaa===="); + + EXPECT_THAT(n, 0); +} + +TEST_F(Base64Test, mju_isValidBase64_invalid7) { + std::size_t n = mju_isValidBase64("A#AA"); + + EXPECT_THAT(n, 0); +} + +TEST_F(Base64Test, mju_isValidBase64_valid1) { + std::size_t n = mju_isValidBase64("AB+/"); + + EXPECT_THAT(n, 3); +} + +TEST_F(Base64Test, mju_isValidBase64_valid2) { + std::size_t n = mju_isValidBase64("ABC="); + + EXPECT_THAT(n, 2); +} + +TEST_F(Base64Test, mju_isValidBase64_valid3) { + std::size_t n = mju_isValidBase64("AB=="); + + EXPECT_THAT(n, 1); +} + +TEST_F(Base64Test, mju_isValidBase64_valid4) { + std::size_t n = mju_isValidBase64("az09AZ+/11=="); + + EXPECT_THAT(n, 7); +} + +TEST_F(Base64Test, mju_decodeBase64) { + std::array buffer; + const char *s = "D4a+//A="; + + std::size_t n = mju_decodeBase64(buffer.data(), s); + + EXPECT_THAT(buffer, ElementsAreArray({15, 134, 190, 255, 240})); + EXPECT_THAT(n, buffer.size()); +} + +TEST_F(Base64Test, mju_decodeBase6_align0) { + std::array buffer; + const char *s = "QUJD"; + + std::size_t n = mju_decodeBase64(buffer.data(), s); + + EXPECT_THAT(buffer, ElementsAreArray({'A', 'B', 'C'})); + EXPECT_THAT(n, buffer.size()); +} + +TEST_F(Base64Test, mju_decodeBase64_align1) { + std::array buffer; + const char *s = "QUI="; + + std::size_t n = mju_decodeBase64(buffer.data(), s); + + EXPECT_THAT(buffer, ElementsAreArray({'A', 'B'})); + EXPECT_THAT(n, buffer.size()); +} + +TEST_F(Base64Test, mju_decodeBase64_align2) { + std::array buffer; + const char *s = "QQ=="; + + std::size_t n = mju_decodeBase64(buffer.data(), s); + + EXPECT_THAT(buffer, ElementsAreArray({'A'})); + EXPECT_THAT(n, buffer.size()); +} + +TEST_F(Base64Test, mju_decodeBase64_null) { + const char *s = ""; + + std::size_t n = mju_decodeBase64(NULL, s); + + EXPECT_THAT(n, 0); +} + +TEST_F(Base64Test, mju_decodeBase64_ones) { + std::array buffer; + const char *s = "////"; + + std::size_t n = mju_decodeBase64(buffer.data(), s); + + EXPECT_THAT(buffer, ElementsAreArray({255, 255, 255})); + EXPECT_THAT(n, buffer.size()); +} + +TEST_F(Base64Test, decodeAndEncode) { + std::array buffer1; + std::array buffer2; + const char *s = "D4a+/vA="; + + mju_decodeBase64(buffer1.data(), s); + mju_encodeBase64(buffer2.data(), buffer1.data(), buffer1.size()); + + EXPECT_THAT(buffer2.data(), StrEq(s)); +} + } // namespace } // namespace mujoco