diff --git a/esphome/components/esphome/ota/__init__.py b/esphome/components/esphome/ota/__init__.py index f5eb878260..cbab552fc2 100644 --- a/esphome/components/esphome/ota/__init__.py +++ b/esphome/components/esphome/ota/__init__.py @@ -283,7 +283,10 @@ FINAL_VALIDATE_SCHEMA = ota_esphome_final_validate FILTER_SOURCE_FILES = filter_source_files_from_defines( - {"ota_esphome_noise.cpp": "USE_OTA_ENCRYPTION"} + { + "ota_esphome_noise.cpp": "USE_OTA_ENCRYPTION", + "ota_esphome_inflate.c": "USE_OTA_DEFLATE", + } ) @@ -305,6 +308,11 @@ async def to_code(config: ConfigType) -> None: if config.get(CONF_ALLOW_PARTITION_ACCESS): cg.add_define("USE_OTA_PARTITIONS") + # ESP8266 inflates gzip in its bootloader; every other platform inflates + # a deflate stream on the fly while it receives the image + if not CORE.is_esp8266: + cg.add_define("USE_OTA_DEFLATE") + # One key per device: an api encryption block supplies it (static or # runtime) and offers; the ota block only adds the requirement api_conf = CORE.config.get(CONF_API) or {} diff --git a/esphome/components/esphome/ota/ota_esphome.cpp b/esphome/components/esphome/ota/ota_esphome.cpp index 1005ed214b..208f1409a6 100644 --- a/esphome/components/esphome/ota/ota_esphome.cpp +++ b/esphome/components/esphome/ota/ota_esphome.cpp @@ -24,6 +24,8 @@ #include #include +#include +#include #include namespace esphome { @@ -179,12 +181,15 @@ static constexpr uint8_t CLIENT_FEATURE_SUPPORTS_COMPRESSION = 0x01; static constexpr uint8_t CLIENT_FEATURE_SUPPORTS_SHA256_AUTH = 0x02; static constexpr uint8_t CLIENT_FEATURE_SUPPORTS_EXTENDED_PROTOCOL = 0x04; static constexpr uint8_t CLIENT_FEATURE_SUPPORTS_NOISE = 0x08; +static constexpr uint8_t CLIENT_FEATURE_SUPPORTS_DEFLATE = 0x10; // Noise needs the extended protocol: the prologue binds the 2-byte feature ack static constexpr uint8_t CLIENT_NOISE_FEATURES = CLIENT_FEATURE_SUPPORTS_NOISE | CLIENT_FEATURE_SUPPORTS_EXTENDED_PROTOCOL; static constexpr uint8_t SERVER_FEATURE_SUPPORTS_COMPRESSION = 0x01; static constexpr uint8_t SERVER_FEATURE_SUPPORTS_PARTITION_ACCESS = 0x02; static constexpr uint8_t SERVER_FEATURE_SUPPORTS_NOISE = 0x04; +// The device inflates a raw deflate stream (window <= OTA_INFLATE_WINDOW_SIZE) +static constexpr uint8_t SERVER_FEATURE_SUPPORTS_DEFLATE = 0x08; inline bool ESPHomeOTAComponent::extended_proto_() const { #ifdef USE_OTA_ENCRYPTION_REQUIRED @@ -310,6 +315,19 @@ void ESPHomeOTAComponent::handle_handshake_() { #elif defined(USE_OTA_ENCRYPTION) // A yaml key always exists: validation rejects the all-zeros key this->handshake_buf_[1] |= SERVER_FEATURE_SUPPORTS_NOISE; +#endif +#ifdef USE_OTA_DEFLATE + // The backend cannot store gzip here (USE_OTA_DEFLATE is not set on the + // one that can), so inflate on the fly when the client offers it and the + // session memory (a few KB) is in hand; otherwise stay uncompressed + if ((this->ota_features_ & CLIENT_FEATURE_SUPPORTS_DEFLATE) != 0 && !supports_compression) { + this->inflate_.reset(new (std::nothrow) InflateSession()); + if (this->inflate_ != nullptr) { + this->handshake_buf_[1] |= SERVER_FEATURE_SUPPORTS_DEFLATE; + } else { + ESP_LOGW(TAG, "No memory to inflate, upload will be uncompressed"); + } + } #endif } else { this->handshake_buf_[0] = @@ -427,16 +445,11 @@ void ESPHomeOTAComponent::handle_data_() { // Backend calls overwrite this with OK; reset to UNKNOWN before any // goto error that follows a successful begin()/write() ota::OTAResponseTypes error_code = ota::OTA_RESPONSE_ERROR_UNKNOWN; - size_t total = 0; - uint32_t last_progress = 0; - uint32_t last_data_ms = 0; + DataTransfer xfer; uint8_t buf[OTA_BUFFER_SIZE]; char *sbuf = reinterpret_cast(buf); - size_t ota_size; + size_t image_size; ota::OTAType ota_type = ota::OTA_TYPE_UPDATE_APP; -#if USE_OTA_VERSION == 2 - size_t size_acknowledged = 0; -#endif // Set socket timeouts and blocking mode (see strategy table above) struct timeval tv; @@ -464,9 +477,20 @@ void ESPHomeOTAComponent::handle_data_() { this->log_read_error_(LOG_STR("size")); goto error; // NOLINT(cppcoreguidelines-avoid-goto) } - ota_size = (static_cast(buf[0]) << 24) | (static_cast(buf[1]) << 16) | - (static_cast(buf[2]) << 8) | buf[3]; - ESP_LOGV(TAG, "Size is %zu bytes", ota_size); + xfer.ota_size = encode_uint32(buf[0], buf[1], buf[2], buf[3]); + ESP_LOGV(TAG, "Size is %zu bytes", xfer.ota_size); + image_size = xfer.ota_size; +#ifdef USE_OTA_DEFLATE + if (this->inflate_ != nullptr) { + // A deflate upload also announces the inflated size, 4 bytes MSB first + if (!this->data_readall_(buf, 4)) { + this->log_read_error_(LOG_STR("image size")); + goto error; // NOLINT(cppcoreguidelines-avoid-goto) + } + image_size = encode_uint32(buf[0], buf[1], buf[2], buf[3]); + ESP_LOGV(TAG, "Inflated size is %zu bytes", image_size); + } +#endif #ifndef USE_OTA_PARTITIONS if (ota_type != ota::OTA_TYPE_UPDATE_APP) { @@ -486,7 +510,7 @@ void ESPHomeOTAComponent::handle_data_() { #endif // begin() returns quickly; flash sectors are erased incrementally during write(). - error_code = this->backend_->begin(ota_size, ota_type); + error_code = this->backend_->begin(image_size, ota_type); if (error_code != ota::OTA_RESPONSE_OK) goto error; // NOLINT(cppcoreguidelines-avoid-goto) @@ -506,76 +530,27 @@ void ESPHomeOTAComponent::handle_data_() { // Acknowledge MD5 OK - 1 byte this->data_write_byte_(ota::OTA_RESPONSE_BIN_MD5_OK); - // Track when we last received data so a silently-vanished peer (no FIN/RST - // delivered, e.g. uploader killed mid-transfer or NAT/router dropped state) - // can't wedge the device indefinitely. Without this, the loop only exits - // on actual data, EOF, or a non-EWOULDBLOCK error from read(), and lwIP - // TCP keepalive isn't enabled here. - last_data_ms = millis(); - while (total < ota_size) { - if (millis() - last_data_ms > OTA_SOCKET_TIMEOUT_DATA) { - ESP_LOGW(TAG, "No data received for %u ms", (unsigned) OTA_SOCKET_TIMEOUT_DATA); - error_code = ota::OTA_RESPONSE_ERROR_UNKNOWN; + xfer.last_data_ms = millis(); +#ifdef USE_OTA_DEFLATE + if (this->inflate_ != nullptr) { + error_code = this->inflate_data_(buf, image_size, xfer); + if (error_code != ota::OTA_RESPONSE_OK) goto error; // NOLINT(cppcoreguidelines-avoid-goto) - } - size_t remaining = ota_size - total; - size_t requested = remaining < OTA_BUFFER_SIZE ? remaining : OTA_BUFFER_SIZE; - ssize_t read; -#ifdef USE_OTA_ENCRYPTION - if (this->noise_ != nullptr) { - // One frame per call; noise_read_data_ waits internally (readall_), so - // there is no would-block retry here and failures are already logged. - read = this->noise_read_data_(buf, requested); - if (read <= 0) { + } else +#endif + { + while (xfer.total < xfer.ota_size) { + ssize_t read = this->receive_data_(buf, xfer); + if (read < 0) { error_code = ota::OTA_RESPONSE_ERROR_UNKNOWN; goto error; // NOLINT(cppcoreguidelines-avoid-goto) } - } else -#endif - { - read = this->client_->read(buf, requested); - if (read == -1) { - const int err = errno; - if (this->would_block_(err)) { - // read() already waited up to SO_RCVTIMEO for data, just feed WDT - App.feed_wdt(); - continue; - } - ESP_LOGW(TAG, "Read err %d", err); - error_code = ota::OTA_RESPONSE_ERROR_UNKNOWN; - goto error; // NOLINT(cppcoreguidelines-avoid-goto) - } else if (read == 0) { - ESP_LOGW(TAG, "Remote closed"); - error_code = ota::OTA_RESPONSE_ERROR_UNKNOWN; + error_code = this->backend_->write(buf, read); + if (error_code != ota::OTA_RESPONSE_OK) { + ESP_LOGW(TAG, "Flash write err %d", error_code); goto error; // NOLINT(cppcoreguidelines-avoid-goto) } } - - last_data_ms = millis(); - error_code = this->backend_->write(buf, read); - if (error_code != ota::OTA_RESPONSE_OK) { - ESP_LOGW(TAG, "Flash write err %d", error_code); - goto error; // NOLINT(cppcoreguidelines-avoid-goto) - } - total += read; -#if USE_OTA_VERSION == 2 - while (size_acknowledged + OTA_BLOCK_SIZE <= total || (total == ota_size && size_acknowledged < ota_size)) { - this->data_write_byte_(ota::OTA_RESPONSE_CHUNK_OK); - size_acknowledged += OTA_BLOCK_SIZE; - } -#endif - - uint32_t now = millis(); - if (now - last_progress > 1000) { - last_progress = now; - float percentage = (total * 100.0f) / ota_size; - ESP_LOGD(TAG, "Progress: %0.1f%%", percentage); -#ifdef USE_OTA_STATE_LISTENER - this->notify_state_(ota::OTA_IN_PROGRESS, percentage, 0); -#endif - // feed watchdog and give other tasks a chance to run - this->yield_and_feed_watchdog_(); - } } // Acknowledge receive OK - 1 byte @@ -771,6 +746,128 @@ bool ESPHomeOTAComponent::try_write_(size_t to_write, const LogString *desc) { return this->handshake_buf_pos_ >= to_write; } +ssize_t ESPHomeOTAComponent::receive_data_(uint8_t *buf, DataTransfer &xfer) { + const size_t remaining = xfer.ota_size - xfer.total; + const size_t requested = std::min(remaining, OTA_BUFFER_SIZE); + ssize_t read; + for (;;) { + // A silently-vanished peer (no FIN/RST delivered, e.g. uploader killed + // mid-transfer or NAT/router dropped state) must not wedge the device: + // read() only fails on EOF or a real error, and lwIP TCP keepalive isn't + // enabled here. + if (millis() - xfer.last_data_ms > OTA_SOCKET_TIMEOUT_DATA) { + ESP_LOGW(TAG, "No data received for %u ms", (unsigned) OTA_SOCKET_TIMEOUT_DATA); + return -1; + } +#ifdef USE_OTA_ENCRYPTION + if (this->noise_ != nullptr) { + // One frame per call; noise_read_data_ waits internally (readall_), so + // there is no would-block retry here and failures are already logged. + read = this->noise_read_data_(buf, requested); + if (read <= 0) + return -1; + break; + } +#endif + read = this->client_->read(buf, requested); + if (read > 0) + break; + if (read == 0) { + ESP_LOGW(TAG, "Remote closed"); + return -1; + } + const int err = errno; + if (!this->would_block_(err)) { + ESP_LOGW(TAG, "Read err %d", err); + return -1; + } + // read() already waited up to SO_RCVTIMEO for data, just feed WDT + App.feed_wdt(); + } + + const uint32_t now = millis(); + xfer.last_data_ms = now; + xfer.total += read; +#if USE_OTA_VERSION == 2 + while (xfer.acknowledged + OTA_BLOCK_SIZE <= xfer.total || + (xfer.total == xfer.ota_size && xfer.acknowledged < xfer.ota_size)) { + this->data_write_byte_(ota::OTA_RESPONSE_CHUNK_OK); + xfer.acknowledged += OTA_BLOCK_SIZE; + } +#endif + if (now - xfer.last_progress > 1000) { + xfer.last_progress = now; + float percentage = (xfer.total * 100.0f) / xfer.ota_size; + ESP_LOGD(TAG, "Progress: %0.1f%%", percentage); +#ifdef USE_OTA_STATE_LISTENER + this->notify_state_(ota::OTA_IN_PROGRESS, percentage, 0); +#endif + // feed watchdog and give other tasks a chance to run + this->yield_and_feed_watchdog_(); + } + return read; +} + +#ifdef USE_OTA_DEFLATE +int ESPHomeOTAComponent::inflate_read_cb_(ota_inflate_state *d) { + // state is the first member, so the session is the same address (checked below) + auto *session = reinterpret_cast(d); + ssize_t read = session->self->receive_data_(session->in, *session->xfer); + if (read <= 0) + return -1; + d->source = session->in + 1; + d->source_limit = session->in + read; + return session->in[0]; +} + +// The window doubles as the output buffer: the decoder fills it, we flush it to +// the backend, and its bytes remain available as the back-reference history for +// the next windowful. +ota::OTAResponseTypes ESPHomeOTAComponent::inflate_data_(uint8_t *in, size_t image_size, DataTransfer &xfer) { + static_assert(offsetof(InflateSession, state) == 0, "inflate_read_cb_ recovers the session from &state"); + InflateSession &session = *this->inflate_; + ota_inflate_state &state = session.state; + session.self = this; + session.xfer = &xfer; + session.in = in; + ota_inflate_init(&state, session.window, OTA_INFLATE_WINDOW_SIZE); + state.source_read_cb = &ESPHomeOTAComponent::inflate_read_cb_; + + size_t written = 0; + int res; + do { + state.dest = session.window; + state.dest_limit = session.window + OTA_INFLATE_WINDOW_SIZE; + res = ota_inflate(&state); + if (res < 0) { + // eof means the read callback failed, which is already logged + if (!state.eof) { + ESP_LOGW(TAG, "Inflate err %d", res); + } + return ota::OTA_RESPONSE_ERROR_UNKNOWN; + } + const size_t produced = state.dest - session.window; + if (produced > image_size - written) { + ESP_LOGW(TAG, "Image exceeds announced size"); + return ota::OTA_RESPONSE_ERROR_UNKNOWN; + } + ota::OTAResponseTypes write_result = this->backend_->write(session.window, produced); + if (write_result != ota::OTA_RESPONSE_OK) { + ESP_LOGW(TAG, "Flash write err %d", write_result); + return write_result; + } + written += produced; + } while (res != OTA_INFLATE_DONE); + + if (written != image_size || xfer.total != xfer.ota_size) { + ESP_LOGW(TAG, "Inflated %zu of %zu bytes from %zu of %zu", written, image_size, xfer.total, xfer.ota_size); + return ota::OTA_RESPONSE_ERROR_UNKNOWN; + } + ESP_LOGD(TAG, "Inflated %zu bytes from %zu", written, xfer.total); + return ota::OTA_RESPONSE_OK; +} +#endif // USE_OTA_DEFLATE + void ESPHomeOTAComponent::cleanup_connection_() { this->client_->close(); this->client_ = nullptr; @@ -784,6 +881,9 @@ void ESPHomeOTAComponent::cleanup_connection_() { #endif #ifdef USE_OTA_ENCRYPTION this->noise_ = nullptr; +#endif +#ifdef USE_OTA_DEFLATE + this->inflate_ = nullptr; #endif // Intentionally no disable_loop() — letting loop() run one more iteration catches // any connection that queued on the listener mid-session (otherwise the wake flag, diff --git a/esphome/components/esphome/ota/ota_esphome.h b/esphome/components/esphome/ota/ota_esphome.h index c6f710b3fc..3187e20e6d 100644 --- a/esphome/components/esphome/ota/ota_esphome.h +++ b/esphome/components/esphome/ota/ota_esphome.h @@ -7,6 +7,9 @@ #ifdef USE_OTA_ENCRYPTION #include "esphome/components/noise/noise_handshake.h" #endif +#ifdef USE_OTA_DEFLATE +#include "ota_esphome_inflate.h" +#endif #include "esphome/core/helpers.h" #include "esphome/core/log.h" #include "esphome/core/preferences.h" @@ -119,6 +122,21 @@ class ESPHomeOTAComponent final : public ota::OTAComponent { return this->readall_(buf, len); } + // Upload accounting shared by the data loop and the inflate read callback + struct DataTransfer { + size_t ota_size; // bytes the client sends + size_t total{0}; // bytes received so far +#if USE_OTA_VERSION == 2 + size_t acknowledged{0}; +#endif + uint32_t last_data_ms; + uint32_t last_progress{0}; + }; + // Receives up to OTA_BUFFER_SIZE bytes of upload data into buf, waiting up to + // the data timeout; updates xfer and sends chunk acks. Returns bytes read, -1 + // on failure (logged). + ssize_t receive_data_(uint8_t *buf, DataTransfer &xfer); + bool try_read_(size_t to_read, const LogString *desc); bool try_write_(size_t to_write, const LogString *desc); @@ -171,6 +189,24 @@ class ESPHomeOTAComponent final : public ota::OTAComponent { static_assert(OTA_BUFFER_SIZE >= NOISE_CLIENT_MAX_PLAINTEXT + noise::MAC_SIZE, "OTA_BUFFER_SIZE must fit a full encrypted data frame"); #endif +#ifdef USE_OTA_DEFLATE + // Deflate back references reach 1 << espota2.DEFLATE_WINDOW_BITS bytes; the + // ring window must be at least that. It also serves as the inflate output + // buffer, so it is flushed to the backend one windowful at a time. + static constexpr size_t OTA_INFLATE_WINDOW_SIZE = 4096; + // Heap-allocated only while a deflate-compressed upload is negotiated. + struct InflateSession { + ota_inflate_state state; // first member: the read callback casts back from it + ESPHomeOTAComponent *self; + DataTransfer *xfer; + uint8_t *in; // caller's buffer for the compressed input, valid during inflate_data_ + uint8_t window[OTA_INFLATE_WINDOW_SIZE]; + }; + static int inflate_read_cb_(ota_inflate_state *d); + ota::OTAResponseTypes inflate_data_(uint8_t *in, size_t image_size, DataTransfer &xfer); + std::unique_ptr inflate_; +#endif + static constexpr uint8_t MAGIC_BYTES[5] = {0x6C, 0x26, 0xF7, 0x5C, 0x45}; // Derived from the feature byte; storing it would pad the trailing bytes bool extended_proto_() const; diff --git a/esphome/components/esphome/ota/ota_esphome_inflate.c b/esphome/components/esphome/ota/ota_esphome_inflate.c new file mode 100644 index 0000000000..8494fb276c --- /dev/null +++ b/esphome/components/esphome/ota/ota_esphome_inflate.c @@ -0,0 +1,494 @@ +/* + * uzlib - tiny deflate/inflate library (deflate, gzip, zlib) + * + * Copyright (c) 2003 by Joergen Ibsen / Jibz + * All Rights Reserved + * http://www.ibsensoftware.com/ + * + * Copyright (c) 2014-2018 by Paul Sokolovsky + * + * This software is provided 'as-is', without any express + * or implied warranty. In no event will the authors be + * held liable for any damages arising from the use of + * this software. + * + * Permission is granted to anyone to use this software + * for any purpose, including commercial applications, + * and to alter it and redistribute it freely, subject to + * the following restrictions: + * + * 1. The origin of this software must not be + * misrepresented; you must not claim that you + * wrote the original software. If you use this + * software in a product, an acknowledgment in + * the product documentation would be appreciated + * but is not required. + * + * 2. Altered source versions must be plainly marked + * as such, and must not be misrepresented as + * being the original software. + * + * 3. This notice may not be removed or altered from + * any source distribution. + */ + +/* + * Altered for ESPHome: this is the raw deflate decoder from uzlib's + * tinflate.c (v2.9.5) with the gzip/zlib header parsers, checksums, + * runtime table builder and in-memory (non ring window) output path + * removed, and the public names prefixed with ota_inflate. + */ + +#include "ota_esphome_inflate.h" + +#define TINF_OK OTA_INFLATE_OK +#define TINF_DONE OTA_INFLATE_DONE +#define TINF_DATA_ERROR OTA_INFLATE_DATA_ERROR +#define TINF_DICT_ERROR OTA_INFLATE_DICT_ERROR +#define TINF_DATA struct ota_inflate_state +#define TINF_TREE ota_inflate_tree_t +#define TINF_ARRAY_SIZE(arr) (sizeof(arr) / sizeof(*(arr))) + +/* every output byte also goes into the ring window */ +#define TINF_PUT(d, c) \ + { \ + *d->dest++ = c; \ + d->dict_ring[d->dict_idx++] = c; \ + if (d->dict_idx == d->dict_size) \ + d->dict_idx = 0; \ + } + +/* --------------------------------------------------- * + * -- uninitialized global data (static structures) -- * + * --------------------------------------------------- */ + +static const unsigned char LENGTH_BITS[30] = {0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 2, 2, + 2, 2, 3, 3, 3, 3, 4, 4, 4, 4, 5, 5, 5, 5}; +static const unsigned short LENGTH_BASE[30] = {3, 4, 5, 6, 7, 8, 9, 10, 11, 13, 15, 17, 19, 23, 27, + 31, 35, 43, 51, 59, 67, 83, 99, 115, 131, 163, 195, 227, 258}; + +static const unsigned char DIST_BITS[30] = {0, 0, 0, 0, 1, 1, 2, 2, 3, 3, 4, 4, 5, 5, 6, + 6, 7, 7, 8, 8, 9, 9, 10, 10, 11, 11, 12, 12, 13, 13}; +static const unsigned short DIST_BASE[30] = {1, 2, 3, 4, 5, 7, 9, 13, 17, 25, + 33, 49, 65, 97, 129, 193, 257, 385, 513, 769, + 1025, 1537, 2049, 3073, 4097, 6145, 8193, 12289, 16385, 24577}; + +/* special ordering of code length codes */ +static const unsigned char CLCIDX[] = {16, 17, 18, 0, 8, 7, 9, 6, 10, 5, 11, 4, 12, 3, 13, 2, 14, 1, 15}; + +/* ----------------------- * + * -- utility functions -- * + * ----------------------- */ + +/* build the fixed huffman trees */ +static void tinf_build_fixed_trees(TINF_TREE *lt, TINF_TREE *dt) { + int i; + + /* build fixed length tree */ + for (i = 0; i < 7; ++i) + lt->table[i] = 0; + + lt->table[7] = 24; + lt->table[8] = 152; + lt->table[9] = 112; + + for (i = 0; i < 24; ++i) + lt->trans[i] = 256 + i; + for (i = 0; i < 144; ++i) + lt->trans[24 + i] = i; + for (i = 0; i < 8; ++i) + lt->trans[24 + 144 + i] = 280 + i; + for (i = 0; i < 112; ++i) + lt->trans[24 + 144 + 8 + i] = 144 + i; + + /* build fixed distance tree */ + for (i = 0; i < 5; ++i) + dt->table[i] = 0; + + dt->table[5] = 32; + + for (i = 0; i < 32; ++i) + dt->trans[i] = i; +} + +/* given an array of code lengths, build a tree */ +static void tinf_build_tree(TINF_TREE *t, const unsigned char *lengths, unsigned int num) { + unsigned short offs[16]; + unsigned int i, sum; + + /* clear code length count table */ + for (i = 0; i < 16; ++i) + t->table[i] = 0; + + /* scan symbol lengths, and sum code length counts */ + for (i = 0; i < num; ++i) + t->table[lengths[i]]++; + + /* In the lengths array, 0 means unused code. So, t->table[0] now contains + number of unused codes. But table's purpose is to contain # of codes of + particular length, and there're 0 codes of length 0. */ + t->table[0] = 0; + + /* compute offset table for distribution sort */ + for (sum = 0, i = 0; i < 16; ++i) { + offs[i] = sum; + sum += t->table[i]; + } + + /* create code->symbol translation table (symbols sorted by code) */ + for (i = 0; i < num; ++i) { + if (lengths[i]) + t->trans[offs[lengths[i]]++] = i; + } +} + +/* ---------------------- * + * -- decode functions -- * + * ---------------------- */ + +static unsigned char uzlib_get_byte(TINF_DATA *d) { + /* If end of source buffer is not reached, return next byte from source + buffer. */ + if (d->source < d->source_limit) { + return *d->source++; + } + + /* Otherwise if there's callback and we haven't seen EOF yet, try to + read next byte using it. (Note: the callback can also update ->source + and ->source_limit). */ + if (d->source_read_cb && !d->eof) { + int val = d->source_read_cb(d); + if (val >= 0) { + return (unsigned char) val; + } + } + + /* Otherwise, we hit EOF (either from ->source_read_cb() or from exhaustion + of the buffer), and it will be "sticky", i.e. further calls to this + function will end up here too. */ + d->eof = true; + + return 0; +} + +/* get one bit from source stream */ +static int tinf_getbit(TINF_DATA *d) { + unsigned int bit; + + /* check if tag is empty */ + if (!d->bitcount--) { + /* load next tag */ + d->tag = uzlib_get_byte(d); + d->bitcount = 7; + } + + /* shift bit out of tag */ + bit = d->tag & 0x01; + d->tag >>= 1; + + return bit; +} + +/* read a num bit value from a stream and add base */ +static unsigned int tinf_read_bits(TINF_DATA *d, int num, int base) { + unsigned int val = 0; + + /* read num bits */ + if (num) { + unsigned int limit = 1 << (num); + unsigned int mask; + + for (mask = 1; mask < limit; mask *= 2) + if (tinf_getbit(d)) + val += mask; + } + + return val + base; +} + +/* given a data stream and a tree, decode a symbol */ +static int tinf_decode_symbol(TINF_DATA *d, TINF_TREE *t) { + int sum = 0, cur = 0, len = 0; + + /* get more bits while code value is above sum */ + do { + cur = 2 * cur + tinf_getbit(d); + + if (++len == TINF_ARRAY_SIZE(t->table)) { + return TINF_DATA_ERROR; + } + + sum += t->table[len]; + cur -= t->table[len]; + + } while (cur >= 0); + + sum += cur; + if (sum < 0 || sum >= TINF_ARRAY_SIZE(t->trans)) { + return TINF_DATA_ERROR; + } + + return t->trans[sum]; +} + +/* given a data stream, decode dynamic trees from it */ +static int tinf_decode_trees(TINF_DATA *d, TINF_TREE *lt, TINF_TREE *dt) { + /* code lengths for 288 literal/len symbols and 32 dist symbols */ + unsigned char lengths[288 + 32]; + unsigned int hlit, hdist, hclen, hlimit; + unsigned int i, num, length; + + /* get 5 bits HLIT (257-286) */ + hlit = tinf_read_bits(d, 5, 257); + + /* get 5 bits HDIST (1-32) */ + hdist = tinf_read_bits(d, 5, 1); + + /* get 4 bits HCLEN (4-19) */ + hclen = tinf_read_bits(d, 4, 4); + + for (i = 0; i < 19; ++i) + lengths[i] = 0; + + /* read code lengths for code length alphabet */ + for (i = 0; i < hclen; ++i) { + /* get 3 bits code length (0-7) */ + unsigned int clen = tinf_read_bits(d, 3, 0); + + lengths[CLCIDX[i]] = clen; + } + + /* build code length tree, temporarily use length tree */ + tinf_build_tree(lt, lengths, 19); + + /* decode code lengths for the dynamic trees */ + hlimit = hlit + hdist; + for (num = 0; num < hlimit;) { + int sym = tinf_decode_symbol(d, lt); + unsigned char fill_value = 0; + int lbits, lbase = 3; + + /* error decoding */ + if (sym < 0) + return sym; + + switch (sym) { + case 16: + /* copy previous code length 3-6 times (read 2 bits) */ + if (num == 0) + return TINF_DATA_ERROR; + fill_value = lengths[num - 1]; + lbits = 2; + break; + case 17: + /* repeat code length 0 for 3-10 times (read 3 bits) */ + lbits = 3; + break; + case 18: + /* repeat code length 0 for 11-138 times (read 7 bits) */ + lbits = 7; + lbase = 11; + break; + default: + /* values 0-15 represent the actual code lengths */ + lengths[num++] = sym; + /* continue the for loop */ + continue; + } + + /* special code length 16-18 are handled here */ + length = tinf_read_bits(d, lbits, lbase); + if (num + length > hlimit) + return TINF_DATA_ERROR; + for (; length; --length) { + lengths[num++] = fill_value; + } + } + + /* Check that there's "end of block" symbol */ + if (lengths[256] == 0) { + return TINF_DATA_ERROR; + } + + /* build dynamic trees */ + tinf_build_tree(lt, lengths, hlit); + tinf_build_tree(dt, lengths + hlit, hdist); + + return TINF_OK; +} + +/* ----------------------------- * + * -- block inflate functions -- * + * ----------------------------- */ + +/* given a stream and two trees, inflate next chunk of output (a byte or more) */ +static int tinf_inflate_block_data(TINF_DATA *d, TINF_TREE *lt, TINF_TREE *dt) { + if (d->curlen == 0) { + unsigned int offs; + int dist; + int sym = tinf_decode_symbol(d, lt); + + if (d->eof) { + return TINF_DATA_ERROR; + } + + /* literal byte */ + if (sym < 256) { + TINF_PUT(d, sym); + return TINF_OK; + } + + /* end of block */ + if (sym == 256) { + return TINF_DONE; + } + + /* substring from sliding dictionary */ + sym -= 257; + if (sym >= 29) { + return TINF_DATA_ERROR; + } + + /* possibly get more bits from length code */ + d->curlen = tinf_read_bits(d, LENGTH_BITS[sym], LENGTH_BASE[sym]); + + dist = tinf_decode_symbol(d, dt); + if (dist >= 30) { + return TINF_DATA_ERROR; + } + + /* possibly get more bits from distance code */ + offs = tinf_read_bits(d, DIST_BITS[dist], DIST_BASE[dist]); + + /* calculate and validate actual LZ offset to use */ + if (offs > d->dict_size) { + return TINF_DICT_ERROR; + } + /* Note: we don't try to catch offset which points to not yet filled + part of the dictionary here. Doing so would require keeping another + variable to track "filled in" size of the dictionary. Appearance of + such an offset cannot lead to accessing memory outside of the + dictionary buffer, and clients which don't want to leak unrelated + information, should explicitly initialize dictionary buffer passed + to uzlib. */ + + d->lz_off = d->dict_idx - offs; + if (d->lz_off < 0) { + d->lz_off += d->dict_size; + } + } + + /* copy next byte from dict substring */ + TINF_PUT(d, d->dict_ring[d->lz_off]); + if ((unsigned) ++d->lz_off == d->dict_size) { + d->lz_off = 0; + } + d->curlen--; + return TINF_OK; +} + +/* inflate next byte from uncompressed block of data */ +static int tinf_inflate_uncompressed_block(TINF_DATA *d) { + if (d->curlen == 0) { + unsigned int length, invlength; + + /* get length */ + length = uzlib_get_byte(d); + length += 256 * uzlib_get_byte(d); + /* get one's complement of length */ + invlength = uzlib_get_byte(d); + invlength += 256 * uzlib_get_byte(d); + /* check length */ + if (length != (~invlength & 0x0000ffff)) + return TINF_DATA_ERROR; + + /* increment length to properly return TINF_DONE below, without + producing data at the same time */ + d->curlen = length + 1; + + /* make sure we start next block on a byte boundary */ + d->bitcount = 0; + } + + if (--d->curlen == 0) { + return TINF_DONE; + } + + unsigned char c = uzlib_get_byte(d); + TINF_PUT(d, c); + return TINF_OK; +} + +/* ---------------------- * + * -- public functions -- * + * ---------------------- */ + +/* initialize decompression structure */ +void ota_inflate_init(TINF_DATA *d, unsigned char *dict, unsigned int dict_len) { + d->eof = 0; + d->bitcount = 0; + d->bfinal = 0; + d->btype = -1; + d->dict_size = dict_len; + d->dict_ring = dict; + d->dict_idx = 0; + d->curlen = 0; +} + +/* inflate next output bytes from compressed stream */ +int ota_inflate(TINF_DATA *d) { + do { + int res; + + /* start a new block */ + if (d->btype == -1) { + int old_btype; + next_blk: + old_btype = d->btype; + /* read final block flag */ + d->bfinal = tinf_getbit(d); + /* read block type (2 bits) */ + d->btype = tinf_read_bits(d, 2, 0); + + if (d->btype == 1 && old_btype != 1) { + /* build fixed huffman trees */ + tinf_build_fixed_trees(&d->ltree, &d->dtree); + } else if (d->btype == 2) { + /* decode trees from stream */ + res = tinf_decode_trees(d, &d->ltree, &d->dtree); + if (res != TINF_OK) { + return res; + } + } + } + + /* process current block */ + switch (d->btype) { + case 0: + /* decompress uncompressed block */ + res = tinf_inflate_uncompressed_block(d); + break; + case 1: + case 2: + /* decompress block with fixed/dynamic huffman trees */ + /* trees were decoded previously, so it's the same routine for both */ + res = tinf_inflate_block_data(d, &d->ltree, &d->dtree); + break; + default: + return TINF_DATA_ERROR; + } + + if (res == TINF_DONE && !d->bfinal) { + /* the block has ended (without producing more data), but we + can't return without data, so start procesing next block */ + goto next_blk; + } + + if (res != TINF_OK) { + return res; + } + + } while (d->dest < d->dest_limit); + + return TINF_OK; +} diff --git a/esphome/components/esphome/ota/ota_esphome_inflate.h b/esphome/components/esphome/ota/ota_esphome_inflate.h new file mode 100644 index 0000000000..61d82e1489 --- /dev/null +++ b/esphome/components/esphome/ota/ota_esphome_inflate.h @@ -0,0 +1,63 @@ +#pragma once +// Raw deflate decoder for compressed OTA uploads, cut down from uzlib +// (https://github.com/pfalcon/uzlib, zlib licence, see the .c file). +// Kept in C so it stays close to upstream; the decoder writes through a +// ring window so the image never has to be held in RAM. + +#include +#include + +#ifdef __cplusplus +extern "C" { +#endif + +enum ota_inflate_result { + OTA_INFLATE_OK = 0, /* more data produced, call again */ + OTA_INFLATE_DONE = 1, /* end of compressed stream reached */ + OTA_INFLATE_DATA_ERROR = -3, + OTA_INFLATE_DICT_ERROR = -5, +}; + +typedef struct { + unsigned short table[16]; /* table of code length counts */ + unsigned short trans[288]; /* code -> symbol translation table */ +} ota_inflate_tree_t; + +struct ota_inflate_state { + /* Next byte in the input buffer and one past its end */ + const unsigned char *source; + const unsigned char *source_limit; + /* Called when source is exhausted; returns the next byte or -1 at EOF. + It may refill source/source_limit for buffered operation. */ + int (*source_read_cb)(struct ota_inflate_state *d); + + unsigned int tag; + unsigned int bitcount; + + /* Output cursor and one past the end of the output buffer */ + unsigned char *dest; + unsigned char *dest_limit; + + bool eof; + + int btype; + int bfinal; + unsigned int curlen; + int lz_off; + /* Ring window holding the last dict_size output bytes for back references */ + unsigned char *dict_ring; + unsigned int dict_size; + unsigned int dict_idx; + + ota_inflate_tree_t ltree; /* dynamic length/symbol tree */ + ota_inflate_tree_t dtree; /* dynamic distance tree */ +}; + +/* dict must be at least as large as the window the encoder used (its max back reference distance) */ +void ota_inflate_init(struct ota_inflate_state *d, unsigned char *dict, unsigned int dict_len); +/* Produce output until dest reaches dest_limit (OK), the stream ends (DONE) or an error occurs */ +int ota_inflate(struct ota_inflate_state *d); + +#ifdef __cplusplus +} +#endif diff --git a/esphome/core/defines.h b/esphome/core/defines.h index eaece6d5ff..4c3fbca6b6 100644 --- a/esphome/core/defines.h +++ b/esphome/core/defines.h @@ -244,6 +244,7 @@ #define USE_RUNTIME_IMAGE_QOI #define USE_RUNTIME_STATS #define USE_OTA +#define USE_OTA_DEFLATE #define USE_OTA_ENCRYPTION #define USE_OTA_ENCRYPTION_FROM_API #define USE_OTA_ENCRYPTION_PROVISIONED diff --git a/esphome/espota2.py b/esphome/espota2.py index ce403c398d..529763aeac 100644 --- a/esphome/espota2.py +++ b/esphome/espota2.py @@ -11,6 +11,7 @@ import secrets import socket import time from typing import Any +import zlib from esphome.core import EsphomeError from esphome.helpers import ProgressBar, resolve_ip_address @@ -65,9 +66,15 @@ CLIENT_FEATURE_SUPPORTS_COMPRESSION = 0x01 CLIENT_FEATURE_SUPPORTS_SHA256_AUTH = 0x02 CLIENT_FEATURE_SUPPORTS_EXTENDED_PROTOCOL = 0x04 CLIENT_FEATURE_SUPPORTS_NOISE = 0x08 +CLIENT_FEATURE_SUPPORTS_DEFLATE = 0x10 SERVER_FEATURE_SUPPORTS_COMPRESSION = 0x01 SERVER_FEATURE_SUPPORTS_PARTITION_ACCESS = 0x02 SERVER_FEATURE_SUPPORTS_NOISE = 0x04 +SERVER_FEATURE_SUPPORTS_DEFLATE = 0x08 + +# Window of the raw deflate stream sent to a device that inflates on the fly; +# the device's OTA_INFLATE_WINDOW_SIZE (4 KB) must be at least 1 << this +DEFLATE_WINDOW_BITS = 12 NOISE_FRAME_INDICATOR = 0x01 NOISE_HANDSHAKE_OK = 0x00 @@ -547,6 +554,7 @@ def perform_ota( CLIENT_FEATURE_SUPPORTS_COMPRESSION | CLIENT_FEATURE_SUPPORTS_SHA256_AUTH | CLIENT_FEATURE_SUPPORTS_EXTENDED_PROTOCOL + | CLIENT_FEATURE_SUPPORTS_DEFLATE ) if noise_psk: features_to_send |= CLIENT_FEATURE_SUPPORTS_NOISE @@ -640,9 +648,17 @@ def perform_ota( f"retry {flag_name}." ) + deflate = False if features & SERVER_FEATURE_SUPPORTS_COMPRESSION: + # The device stores the gzip file and inflates it when it reboots upload_contents = gzip.compress(file_contents, compresslevel=9) _LOGGER.info("Compressed to %s bytes", len(upload_contents)) + elif extended_proto and features & SERVER_FEATURE_SUPPORTS_DEFLATE: + # The device inflates while receiving through a small ring window + compressor = zlib.compressobj(9, zlib.DEFLATED, -DEFLATE_WINDOW_BITS) + upload_contents = compressor.compress(file_contents) + compressor.flush() + deflate = True + _LOGGER.info("Compressed to %s bytes (deflate)", len(upload_contents)) else: upload_contents = file_contents @@ -701,22 +717,23 @@ def perform_ota( send_check(sock, ota_type, "ota type") upload_size = len(upload_contents) - upload_size_encoded = [ - (upload_size >> 24) & 0xFF, - (upload_size >> 16) & 0xFF, - (upload_size >> 8) & 0xFF, - (upload_size >> 0) & 0xFF, - ] + upload_size_encoded = upload_size.to_bytes(4, "big") # The device erases flash between receiving the size and acking the # prepare, so this window shows the erase cost (near zero when the # device erases lazily during the upload) prepare_start = time.perf_counter() send_check(sock, upload_size_encoded, "binary size") + if deflate: + # The device sizes the partition by the inflated image; its own frame, + # as an encrypted session carries one field per frame + send_check(sock, len(file_contents).to_bytes(4, "big"), "image size") receive_exactly(sock, 1, "update prepare result", RESPONSE_UPDATE_PREPARE_OK) prepare_duration = time.perf_counter() - prepare_start _LOGGER.info("Preparing for upload took %.2f seconds", prepare_duration) - upload_md5 = hashlib.md5(upload_contents).hexdigest() + # The device hashes what it writes to flash: the inflated image for a + # deflate upload, the received bytes otherwise (the gzip file on ESP8266) + upload_md5 = hashlib.md5(file_contents if deflate else upload_contents).hexdigest() _LOGGER.debug("MD5 of upload is %s", upload_md5) send_check(sock, upload_md5, "file checksum") diff --git a/script/ci-custom.py b/script/ci-custom.py index e2b7cd8d37..4f30ba945d 100755 --- a/script/ci-custom.py +++ b/script/ci-custom.py @@ -823,6 +823,8 @@ def lint_relative_py_import(fname: Path, line, col, content): # neither can live in a C++ namespace. "esphome/components/esp32_hosted/esp_now_hosted.cpp", "esphome/components/esp32_hosted/esp_now_hosted_rpc.h", + # C header shared with the vendored decoder + "esphome/components/esphome/ota/ota_esphome_inflate.h", ], ) def lint_namespace(fname: Path, content: str) -> str | None: diff --git a/tests/integration/test_host_ota.py b/tests/integration/test_host_ota.py index f8c122c6e1..75a779e6b2 100644 --- a/tests/integration/test_host_ota.py +++ b/tests/integration/test_host_ota.py @@ -180,10 +180,14 @@ async def test_host_ota_self_update( ) ) staged = asyncio.Event() + inflated = asyncio.Event() def on_log(line: str) -> None: if "OTA staged at" in line: staged.set() + # The host backend has no gzip support, so the upload negotiates deflate + if "Inflated " in line: + inflated.set() dev.on_log(line) async with run_binary(dev.binary_path, line_callback=on_log) as (proc, _lines): @@ -195,6 +199,7 @@ async def test_host_ota_self_update( await dev.ota(None, None, "espota2 reported failure") assert staged.is_set() + assert inflated.is_set(), "upload was not deflate compressed" async with wait_and_connect_api_client(port=dev.api_port) as client: info_after = await client.device_info() diff --git a/tests/unit_tests/test_espota2.py b/tests/unit_tests/test_espota2.py index 8867e2c215..54112e4d2b 100644 --- a/tests/unit_tests/test_espota2.py +++ b/tests/unit_tests/test_espota2.py @@ -12,6 +12,7 @@ from pathlib import Path import socket import struct from unittest.mock import Mock, call, patch +import zlib import pytest from pytest import CaptureFixture @@ -354,6 +355,7 @@ def test_perform_ota_successful_md5_auth( espota2.CLIENT_FEATURE_SUPPORTS_COMPRESSION | espota2.CLIENT_FEATURE_SUPPORTS_SHA256_AUTH | espota2.CLIENT_FEATURE_SUPPORTS_EXTENDED_PROTOCOL + | espota2.CLIENT_FEATURE_SUPPORTS_DEFLATE ] ) ) @@ -1051,6 +1053,7 @@ def test_perform_ota_successful_sha256_auth( espota2.CLIENT_FEATURE_SUPPORTS_COMPRESSION | espota2.CLIENT_FEATURE_SUPPORTS_SHA256_AUTH | espota2.CLIENT_FEATURE_SUPPORTS_EXTENDED_PROTOCOL + | espota2.CLIENT_FEATURE_SUPPORTS_DEFLATE ] ) ) @@ -1107,6 +1110,7 @@ def test_perform_ota_sha256_fallback_to_md5( espota2.CLIENT_FEATURE_SUPPORTS_COMPRESSION | espota2.CLIENT_FEATURE_SUPPORTS_SHA256_AUTH | espota2.CLIENT_FEATURE_SUPPORTS_EXTENDED_PROTOCOL + | espota2.CLIENT_FEATURE_SUPPORTS_DEFLATE ] ) ) @@ -1216,6 +1220,7 @@ def test_perform_ota_extended_protocol_app( espota2.CLIENT_FEATURE_SUPPORTS_COMPRESSION | espota2.CLIENT_FEATURE_SUPPORTS_SHA256_AUTH | espota2.CLIENT_FEATURE_SUPPORTS_EXTENDED_PROTOCOL + | espota2.CLIENT_FEATURE_SUPPORTS_DEFLATE ] ) ) @@ -1276,6 +1281,7 @@ def test_perform_ota_successful_partition_table( espota2.CLIENT_FEATURE_SUPPORTS_COMPRESSION | espota2.CLIENT_FEATURE_SUPPORTS_SHA256_AUTH | espota2.CLIENT_FEATURE_SUPPORTS_EXTENDED_PROTOCOL + | espota2.CLIENT_FEATURE_SUPPORTS_DEFLATE ] ) ) @@ -1504,3 +1510,57 @@ def test_check_error_passes_non_error_when_expect_is_none() -> None: espota2.check_error([espota2.RESPONSE_OK], None) espota2.check_error([espota2.RESPONSE_HEADER_OK], None) espota2.check_error([espota2.RESPONSE_FEATURE_FLAGS], None) + + +def _deflate_handshake(server_features: int) -> list[bytes]: + return [ + bytes([espota2.RESPONSE_OK]), + bytes([espota2.OTA_VERSION_2_0]), + bytes([espota2.RESPONSE_FEATURE_FLAGS]), + bytes([server_features]), + bytes([espota2.RESPONSE_AUTH_OK]), + bytes([espota2.RESPONSE_UPDATE_PREPARE_OK]), + bytes([espota2.RESPONSE_BIN_MD5_OK]), + bytes([espota2.RESPONSE_CHUNK_OK]), + bytes([espota2.RESPONSE_RECEIVE_OK]), + bytes([espota2.RESPONSE_UPDATE_END_OK]), + ] + + +@pytest.mark.usefixtures("mock_time") +def test_perform_ota_with_deflate(mock_socket: Mock) -> None: + """A device that inflates on the fly gets a raw deflate stream, both sizes and the image MD5.""" + original_content = b"firmware" * 100 + mock_socket.recv.side_effect = _deflate_handshake( + espota2.SERVER_FEATURE_SUPPORTS_DEFLATE + ) + + espota2.perform_ota(mock_socket, None, io.BytesIO(original_content), "test.bin") + + sent = [c[0][0] for c in mock_socket.sendall.call_args_list] + # magic, features, ota type, size, image size, md5, data, end ack + sent_size = struct.unpack(">I", sent[3])[0] + assert sent[4] == len(original_content).to_bytes(4, "big") + payload = sent[6] + assert len(payload) == sent_size < len(original_content) + # The device decodes through a window of 1 << DEFLATE_WINDOW_BITS bytes + assert zlib.decompress(payload, -espota2.DEFLATE_WINDOW_BITS) == original_content + assert sent[5] == hashlib.md5(original_content).hexdigest().encode() + + +@pytest.mark.usefixtures("mock_time") +def test_perform_ota_gzip_wins_over_deflate(mock_socket: Mock) -> None: + """A device that can store gzip keeps getting gzip even when it also offers deflate.""" + original_content = b"firmware" * 100 + mock_socket.recv.side_effect = _deflate_handshake( + espota2.SERVER_FEATURE_SUPPORTS_COMPRESSION + | espota2.SERVER_FEATURE_SUPPORTS_DEFLATE + ) + + espota2.perform_ota(mock_socket, None, io.BytesIO(original_content), "test.bin") + + sent = [c[0][0] for c in mock_socket.sendall.call_args_list] + compressed = gzip.compress(original_content, compresslevel=9) + assert sent[3] == len(compressed).to_bytes(4, "big") + assert sent[4] == hashlib.md5(compressed).hexdigest().encode() + assert sent[5] == compressed