Add Base64, defined by RFC 4648, encoding and decoding to MuJoCo engine.
PiperOrigin-RevId: 561008882 Change-Id: I494eb163bc82e8a5c9ad2116ef119072b08a9f48
This commit is contained in:
committed by
Copybara-Service
parent
50e968a6e2
commit
389736f5a1
@@ -14,7 +14,9 @@
|
||||
|
||||
#include "engine/engine_util_misc.h"
|
||||
|
||||
#include <ctype.h>
|
||||
#include <math.h>
|
||||
#include <stdint.h>
|
||||
#include <stdio.h>
|
||||
#include <stdlib.h>
|
||||
#include <string.h>
|
||||
@@ -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++) {
|
||||
|
||||
@@ -17,12 +17,14 @@
|
||||
|
||||
#include <mujoco/mjexport.h>
|
||||
#include <mujoco/mjmodel.h>
|
||||
#include <mujoco/mjtnum.h>
|
||||
|
||||
#ifdef __cplusplus
|
||||
extern "C" {
|
||||
#endif
|
||||
|
||||
#include <stddef.h>
|
||||
#include <stdint.h>
|
||||
|
||||
//------------------------------ 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
|
||||
|
||||
@@ -14,6 +14,11 @@
|
||||
|
||||
// Tests for engine/engine_util_solve.c.
|
||||
|
||||
#include <array>
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <cstring>
|
||||
|
||||
#include <gmock/gmock.h>
|
||||
#include <gtest/gtest.h>
|
||||
#include <mujoco/mjdata.h>
|
||||
@@ -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<char, 9> buffer;
|
||||
std::array<std::uint8_t, 5> 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<char, 5> buffer;
|
||||
std::array<std::uint8_t, 3> 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<char, 5> buffer;
|
||||
std::array<std::uint8_t, 2> 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<char, 5> buffer;
|
||||
std::array<std::uint8_t, 1> 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<char, 1> 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<char, 5> buffer;
|
||||
std::array<std::uint8_t, 3> 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<std::uint8_t, 5> 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<std::uint8_t, 3> 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<std::uint8_t, 2> 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<std::uint8_t, 1> 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<std::uint8_t, 3> 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<std::uint8_t, 5> buffer1;
|
||||
std::array<char, 9> 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
|
||||
|
||||
Reference in New Issue
Block a user