mirror of
https://github.com/esphome/esphome.git
synced 2026-10-07 11:26:39 +00:00
[core] Decode base64 in a single pass (#20123)
This commit is contained in:
+22
-62
@@ -8,7 +8,6 @@
|
||||
|
||||
#include <strings.h>
|
||||
#include <algorithm>
|
||||
#include <cctype>
|
||||
#include <cmath>
|
||||
#include <cstdarg>
|
||||
#include <cstdio>
|
||||
@@ -621,15 +620,11 @@ int8_t step_to_accuracy_decimals(float step) {
|
||||
return decimals;
|
||||
}
|
||||
|
||||
// Map a base64/base64url character to its 6-bit value (0-63) arithmetically.
|
||||
// No lookup table: a table would occupy RAM on ESP8266 (.rodata lives in DRAM there).
|
||||
// Supports both standard base64 (+/) and base64url (-_) alphabets.
|
||||
// NOTE: This returns 0 for both 'A' (valid base64 char at index 0) and invalid characters.
|
||||
// This is safe because is_base64() is ALWAYS checked before calling this function,
|
||||
// preventing invalid characters from ever reaching here. The base64_decode function
|
||||
// stops processing at the first invalid character due to the is_base64() check in its
|
||||
// while loop condition, making this edge case harmless in practice.
|
||||
static inline uint8_t base64_find_char(char c) {
|
||||
static constexpr uint8_t INVALID_BASE64_CHAR = 0xFF;
|
||||
|
||||
// 6-bit value of a base64 or base64url char, or INVALID_BASE64_CHAR.
|
||||
// No lookup table: .rodata lives in DRAM on ESP8266.
|
||||
static constexpr uint8_t base64_char_value(uint8_t c) {
|
||||
if (c >= 'A' && c <= 'Z')
|
||||
return c - 'A';
|
||||
if (c >= 'a' && c <= 'z')
|
||||
@@ -641,70 +636,35 @@ static inline uint8_t base64_find_char(char c) {
|
||||
return 62;
|
||||
if (c == '/' || c == '_')
|
||||
return 63;
|
||||
return 0;
|
||||
return INVALID_BASE64_CHAR;
|
||||
}
|
||||
|
||||
// Check if character is valid base64 or base64url
|
||||
static inline bool is_base64(char c) { return (isalnum(c) || (c == '+') || (c == '/') || (c == '-') || (c == '_')); }
|
||||
|
||||
// base64_encode (both overloads) moved to alloc_helpers.cpp
|
||||
|
||||
size_t base64_decode(const std::string &encoded_string, uint8_t *buf, size_t buf_len) {
|
||||
return base64_decode(reinterpret_cast<const uint8_t *>(encoded_string.data()), encoded_string.size(), buf, buf_len);
|
||||
}
|
||||
|
||||
// Decode 4 base64 characters to up to 'count' output bytes, returns true if truncated.
|
||||
static inline bool base64_decode_quad(uint8_t *char_array_4, int count, uint8_t *buf, size_t buf_len, size_t &out) {
|
||||
for (int i = 0; i < 4; i++)
|
||||
char_array_4[i] = base64_find_char(char_array_4[i]);
|
||||
|
||||
uint8_t char_array_3[3];
|
||||
char_array_3[0] = (char_array_4[0] << 2) + ((char_array_4[1] & 0x30) >> 4);
|
||||
char_array_3[1] = ((char_array_4[1] & 0xf) << 4) + ((char_array_4[2] & 0x3c) >> 2);
|
||||
char_array_3[2] = ((char_array_4[2] & 0x3) << 6) + char_array_4[3];
|
||||
|
||||
bool truncated = false;
|
||||
for (int j = 0; j < count; j++) {
|
||||
if (out < buf_len) {
|
||||
buf[out++] = char_array_3[j];
|
||||
} else {
|
||||
truncated = true;
|
||||
}
|
||||
}
|
||||
return truncated;
|
||||
}
|
||||
|
||||
size_t base64_decode(const uint8_t *encoded_data, size_t encoded_len, uint8_t *buf, size_t buf_len) {
|
||||
size_t in_len = encoded_len;
|
||||
int i = 0;
|
||||
size_t in = 0;
|
||||
size_t out = 0;
|
||||
uint8_t char_array_4[4];
|
||||
bool truncated = false;
|
||||
|
||||
// SAFETY: The loop condition checks is_base64() before processing each character.
|
||||
// This ensures base64_find_char() is only called on valid base64 characters,
|
||||
// preventing the edge case where invalid chars would return 0 (same as 'A').
|
||||
while (in_len-- && (encoded_data[in] != '=') && is_base64(encoded_data[in])) {
|
||||
char_array_4[i++] = encoded_data[in];
|
||||
in++;
|
||||
if (i == 4) {
|
||||
truncated |= base64_decode_quad(char_array_4, 3, buf, buf_len, out);
|
||||
i = 0;
|
||||
uint32_t accum = 0;
|
||||
uint32_t bits = 0;
|
||||
// Stops at '=' or any non-alphabet char; leftover bits of a partial group are dropped.
|
||||
for (size_t in = 0; in < encoded_len; in++) {
|
||||
uint8_t value = base64_char_value(encoded_data[in]);
|
||||
if (value == INVALID_BASE64_CHAR)
|
||||
break;
|
||||
accum = (accum << 6) | value;
|
||||
bits += 6;
|
||||
if (bits >= 8) {
|
||||
bits -= 8;
|
||||
if (out == buf_len) {
|
||||
ESP_LOGW(TAG, "Base64 decode: buffer too small, truncating");
|
||||
return out;
|
||||
}
|
||||
buf[out++] = static_cast<uint8_t>(accum >> bits);
|
||||
}
|
||||
}
|
||||
|
||||
if (i) {
|
||||
for (int j = i; j < 4; j++)
|
||||
char_array_4[j] = 0;
|
||||
|
||||
truncated |= base64_decode_quad(char_array_4, i - 1, buf, buf_len, out);
|
||||
}
|
||||
|
||||
if (truncated) {
|
||||
ESP_LOGW(TAG, "Base64 decode: buffer too small, truncating");
|
||||
}
|
||||
|
||||
return out;
|
||||
}
|
||||
|
||||
|
||||
@@ -281,6 +281,27 @@ TEST(Base64, Rfc4648Vectors) {
|
||||
}
|
||||
}
|
||||
|
||||
TEST(Base64, DecodeTruncatesToBuffer) {
|
||||
uint8_t buf[4];
|
||||
size_t len = base64_decode(std::string("Zm9vYmFy"), buf, sizeof(buf));
|
||||
EXPECT_EQ(len, 4u);
|
||||
EXPECT_EQ(memcmp(buf, "foob", 4), 0);
|
||||
}
|
||||
|
||||
TEST(Base64, DecodeStopsAtNonAlphabetChar) {
|
||||
uint8_t buf[8];
|
||||
EXPECT_EQ(base64_decode(std::string("Zm9v!Zm9v"), buf, sizeof(buf)), 3u);
|
||||
EXPECT_EQ(memcmp(buf, "foo", 3), 0);
|
||||
EXPECT_EQ(base64_decode(std::string("Zm9v Zm9v"), buf, sizeof(buf)), 3u);
|
||||
EXPECT_EQ(base64_decode(std::string("Zm9v\xC3Zm9v"), buf, sizeof(buf)), 3u);
|
||||
}
|
||||
|
||||
TEST(Base64, DecodeDropsPartialGroup) {
|
||||
uint8_t buf[8];
|
||||
EXPECT_EQ(base64_decode(std::string("Z"), buf, sizeof(buf)), 0u);
|
||||
EXPECT_EQ(base64_decode(std::string("Zm9vY"), buf, sizeof(buf)), 3u);
|
||||
}
|
||||
|
||||
// --- step_to_accuracy_decimals() ---
|
||||
|
||||
TEST(StepToAccuracyDecimals, TypicalSteps) {
|
||||
|
||||
Reference in New Issue
Block a user