mirror of
https://github.com/esphome/esphome.git
synced 2026-09-10 14:57:32 +00:00
[esphome.ota] Compress OTA uploads with deflate on platforms without gzip support
ESP32, RP2040, LibreTiny and host could not receive a compressed image; only the ESP8266 can, because its bootloader inflates a gzip file at reboot. This inflates a raw deflate stream on the fly through a 4 KB ring window that also serves as the output buffer, so the device never holds the whole image. The CLI offers a new client feature bit; a device whose backend has no gzip support and has the inflater compiled answers with a new server bit once the session memory is in hand, then the CLI sends a 4 KB window deflate stream and the MD5 of the inflated image. On allocation failure the device declines the bit and the upload stays uncompressed. Old CLIs and old devices never set the bits, so both directions stay compatible; the ESP8266 keeps its gzip path and the CLI prefers gzip when a device offers both. The decoder is uzlib's tinflate.c (zlib licence) trimmed to raw deflate.
This commit is contained in:
@@ -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 {}
|
||||
|
||||
@@ -24,6 +24,8 @@
|
||||
|
||||
#include <cerrno>
|
||||
#include <cstdio>
|
||||
#include <cstddef>
|
||||
#include <new>
|
||||
#include <sys/time.h>
|
||||
|
||||
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<char *>(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<size_t>(buf[0]) << 24) | (static_cast<size_t>(buf[1]) << 16) |
|
||||
(static_cast<size_t>(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<InflateSession *>(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,
|
||||
|
||||
@@ -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<InflateSession> 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;
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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 <stdbool.h>
|
||||
#include <stdint.h>
|
||||
|
||||
#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
|
||||
@@ -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
|
||||
|
||||
+24
-7
@@ -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")
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user