diff --git a/esphome/components/api/api_server.h b/esphome/components/api/api_server.h index 072a583901..146ff09381 100644 --- a/esphome/components/api/api_server.h +++ b/esphome/components/api/api_server.h @@ -78,7 +78,7 @@ class APIServer final : public Component, #ifdef USE_API_NOISE bool save_noise_psk(noise::psk_t psk, bool make_active = true); bool clear_noise_psk(bool make_active = true); - void set_noise_psk(noise::psk_t psk) { this->noise_ctx_.set_psk(psk); } + void set_noise_psk(const noise::psk_t &psk) { this->noise_ctx_.set_psk(psk); } noise::NoiseContext &get_noise_ctx() { return this->noise_ctx_; } #endif // USE_API_NOISE diff --git a/esphome/components/esphome/ota/ota_esphome.h b/esphome/components/esphome/ota/ota_esphome.h index fd164b8138..6c75980525 100644 --- a/esphome/components/esphome/ota/ota_esphome.h +++ b/esphome/components/esphome/ota/ota_esphome.h @@ -45,7 +45,7 @@ class ESPHomeOTAComponent final : public ota::OTAComponent { #endif // USE_OTA_PASSWORD #ifdef USE_OTA_ENCRYPTION - void set_noise_psk(noise::psk_t psk) { this->noise_ctx_.set_psk(psk); } + void set_noise_psk(const noise::psk_t &psk) { this->noise_ctx_.set_psk(psk); } #endif /// Manually set the port OTA should listen on @@ -88,6 +88,7 @@ class ESPHomeOTAComponent final : public ota::OTAComponent { bool noise_start_session_(uint8_t server_feature_flags); bool handle_noise_handshake_(); bool noise_try_read_frame_(); + size_t noise_frame_payload_len_(const uint8_t *header, size_t min_len, size_t max_len); bool noise_try_write_frame_(); void noise_send_reject_(const LogString *reason); ssize_t noise_decrypt_(uint8_t *buf, size_t len); diff --git a/esphome/components/esphome/ota/ota_esphome_noise.cpp b/esphome/components/esphome/ota/ota_esphome_noise.cpp index 847b3e95b7..7dcc0ec5fd 100644 --- a/esphome/components/esphome/ota/ota_esphome_noise.cpp +++ b/esphome/components/esphome/ota/ota_esphome_noise.cpp @@ -42,12 +42,6 @@ ESPHomeOTAComponent::NoiseSession::~NoiseSession() { bool ESPHomeOTAComponent::noise_start_session_(uint8_t server_feature_flags) { // NOLINTNEXTLINE(clang-analyzer-cplusplus.NewDeleteLeaks) this->noise_ = std::unique_ptr(new (std::nothrow) NoiseSession()); - if (this->noise_ == nullptr) { - ESP_LOGW(TAG, "Session allocation failed"); - this->cleanup_connection_(); - return false; - } - static constexpr size_t PROLOGUE_ACK_LEN = 2; // OTA_RESPONSE_OK + version static constexpr size_t PROLOGUE_CLIENT_FEATURES_LEN = 1; static constexpr size_t PROLOGUE_FEATURE_ACK_LEN = 2; // OTA_RESPONSE_FEATURE_FLAGS + server flags @@ -71,9 +65,11 @@ bool ESPHomeOTAComponent::noise_start_session_(uint8_t server_feature_flags) { *p++ = ota::OTA_RESPONSE_FEATURE_FLAGS; *p++ = server_feature_flags; - int err = this->noise_->handshake.init(this->noise_ctx_.get_psk(), prologue, sizeof(prologue)); + int err = this->noise_ == nullptr + ? NOISE_ERROR_NO_MEMORY + : this->noise_->handshake.init(this->noise_ctx_.get_psk(), prologue, sizeof(prologue)); if (err != 0) { - ESP_LOGW(TAG, "Handshake init: %d", err); + ESP_LOGW(TAG, "Session init: %d", err); this->cleanup_connection_(); return false; } @@ -156,20 +152,30 @@ bool ESPHomeOTAComponent::handle_noise_handshake_() { } } +/// Payload length from a frame header, or 0 (logged) when the indicator or +/// the length is out of range. +size_t ESPHomeOTAComponent::noise_frame_payload_len_(const uint8_t *header, size_t min_len, size_t max_len) { + const size_t payload_len = encode_uint16(header[1], header[2]); + if (header[0] != noise::FRAME_INDICATOR || payload_len < min_len || payload_len > max_len) { + ESP_LOGW(TAG, "Bad frame: 0x%02X, %zu bytes", header[0], payload_len); + return 0; + } + return payload_len; +} + /// Non-blocking read of one handshake frame into the session buffer. bool ESPHomeOTAComponent::noise_try_read_frame_() { NoiseSession &s = *this->noise_; while (s.frame_pos < noise::FRAME_HEADER_SIZE) { ssize_t read = this->client_->read(s.frame_buf + s.frame_pos, noise::FRAME_HEADER_SIZE - s.frame_pos); - if (!this->handle_read_error_(read, LOG_STR("read noise frame"))) { + if (!this->handle_read_error_(read, LOG_STR("read noise"))) { return false; } s.frame_pos += read; } if (s.frame_len == 0) { - const uint16_t payload_len = encode_uint16(s.frame_buf[1], s.frame_buf[2]); - if (s.frame_buf[0] != noise::FRAME_INDICATOR || payload_len < 1 || payload_len > 1 + noise::MAX_HANDSHAKE_SIZE) { - ESP_LOGW(TAG, "Bad handshake frame: 0x%02X, %u bytes", s.frame_buf[0], payload_len); + const size_t payload_len = this->noise_frame_payload_len_(s.frame_buf, 1, 1 + noise::MAX_HANDSHAKE_SIZE); + if (payload_len == 0) { this->cleanup_connection_(); return false; } @@ -177,7 +183,7 @@ bool ESPHomeOTAComponent::noise_try_read_frame_() { } while (s.frame_pos < s.frame_len) { ssize_t read = this->client_->read(s.frame_buf + s.frame_pos, s.frame_len - s.frame_pos); - if (!this->handle_read_error_(read, LOG_STR("read noise frame"))) { + if (!this->handle_read_error_(read, LOG_STR("read noise"))) { return false; } s.frame_pos += read; @@ -231,9 +237,8 @@ ssize_t ESPHomeOTAComponent::noise_read_frame_blocking_(uint8_t *buf, size_t min if (!this->readall_(header, sizeof(header))) { return -1; } - const size_t ciphertext_len = encode_uint16(header[1], header[2]); - if (header[0] != noise::FRAME_INDICATOR || ciphertext_len < min_ciphertext || ciphertext_len > max_ciphertext) { - ESP_LOGW(TAG, "Bad frame: 0x%02X, %zu bytes", header[0], ciphertext_len); + const size_t ciphertext_len = this->noise_frame_payload_len_(header, min_ciphertext, max_ciphertext); + if (ciphertext_len == 0) { return -1; } if (!this->readall_(buf, ciphertext_len)) { diff --git a/esphome/components/noise/noise.cpp b/esphome/components/noise/noise.cpp index 95fab322db..7c79e98098 100644 --- a/esphome/components/noise/noise.cpp +++ b/esphome/components/noise/noise.cpp @@ -15,6 +15,11 @@ namespace esphome::noise { static const char *const TAG = "noise"; +void NoiseContext::set_psk(const psk_t &psk) { + this->psk_ = psk; + this->has_psk_ = !is_all_zeros(psk); +} + const LogString *noise_err_to_logstr(int err) { if (err == NOISE_ERROR_NO_MEMORY) return LOG_STR("NO_MEMORY"); diff --git a/esphome/components/noise/noise.h b/esphome/components/noise/noise.h index f9da8d35b8..1e779f0380 100644 --- a/esphome/components/noise/noise.h +++ b/esphome/components/noise/noise.h @@ -23,10 +23,7 @@ class NoiseContext { } return acc == 0; } - void set_psk(psk_t psk) { - this->psk_ = psk; - this->has_psk_ = !is_all_zeros(psk); - } + void set_psk(const psk_t &psk); const psk_t &get_psk() const { return this->psk_; } bool has_psk() const { return this->has_psk_; }