From f334b9ea1106a289b44439dc1c139db24e4e21fa Mon Sep 17 00:00:00 2001 From: "J. Nick Koston" Date: Mon, 5 Oct 2026 10:54:17 -0500 Subject: [PATCH] [core] Decode base64 in a single pass (#20123) --- esphome/core/helpers.cpp | 84 +++++++------------------- tests/components/core/test_helpers.cpp | 21 +++++++ 2 files changed, 43 insertions(+), 62 deletions(-) diff --git a/esphome/core/helpers.cpp b/esphome/core/helpers.cpp index 313daacbd0..a11b07808d 100644 --- a/esphome/core/helpers.cpp +++ b/esphome/core/helpers.cpp @@ -8,7 +8,6 @@ #include #include -#include #include #include #include @@ -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(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(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; } diff --git a/tests/components/core/test_helpers.cpp b/tests/components/core/test_helpers.cpp index 72af605d61..9fee0c723b 100644 --- a/tests/components/core/test_helpers.cpp +++ b/tests/components/core/test_helpers.cpp @@ -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) {