diff --git a/esphome/components/esphome/ota/ota_esphome.cpp b/esphome/components/esphome/ota/ota_esphome.cpp index d353d01d20..6248474624 100644 --- a/esphome/components/esphome/ota/ota_esphome.cpp +++ b/esphome/components/esphome/ota/ota_esphome.cpp @@ -467,6 +467,8 @@ void ESPHomeOTAComponent::handle_data_() { if (this->extended_proto_()) { // Read ota type, 1 byte if (!this->data_readall_(buf, 1)) { + if (this->client_left_before_start_()) + return; this->log_read_error_(LOG_STR("OTA type")); goto error; // NOLINT(cppcoreguidelines-avoid-goto) } @@ -476,6 +478,9 @@ void ESPHomeOTAComponent::handle_data_() { // Read size, 4 bytes MSB first if (!this->data_readall_(buf, 4)) { + // The first request byte is the type on the extended protocol; a close after it was a cut-off request + if (!this->extended_proto_() && this->client_left_before_start_()) + return; this->log_read_error_(LOG_STR("size")); goto error; // NOLINT(cppcoreguidelines-avoid-goto) } @@ -542,6 +547,8 @@ void ESPHomeOTAComponent::handle_data_() { // there is no would-block retry here and failures are already logged. read = this->noise_read_data_(buf, requested); if (read <= 0) { + if (this->remote_closed_) + this->log_remote_closed_(LOG_STR("data")); error_code = ota::OTA_RESPONSE_ERROR_UNKNOWN; goto error; // NOLINT(cppcoreguidelines-avoid-goto) } @@ -666,7 +673,11 @@ bool ESPHomeOTAComponent::readall_(uint8_t *buf, size_t len) { return false; } } else if (read == 0) { - ESP_LOGW(TAG, "Remote closed"); + // A partial message is a cut-off request, not a clean close; the caller reports the clean one + this->remote_closed_ = at == 0; + if (at > 0) { + ESP_LOGW(TAG, "Remote closed after %u of %zu bytes", (unsigned) at, len); + } return false; } else { at += read; @@ -712,7 +723,22 @@ void ESPHomeOTAComponent::log_socket_error_(const LogString *msg) { ESP_LOGW(TAG, "Socket %s: errno %d", LOG_STR_ARG(msg), errno); } -void ESPHomeOTAComponent::log_read_error_(const LogString *what) { ESP_LOGW(TAG, "Read %s failed", LOG_STR_ARG(what)); } +bool ESPHomeOTAComponent::client_left_before_start_() { + // Key probes and scanners hang up right after the handshake; nothing started, so no error status or callback + if (!this->remote_closed_) + return false; + ESP_LOGD(TAG, "Client left after the handshake"); + this->cleanup_connection_(); + return true; +} + +void ESPHomeOTAComponent::log_read_error_(const LogString *what) { + if (this->remote_closed_) { + this->log_remote_closed_(what); + return; + } + ESP_LOGW(TAG, "Read %s failed", LOG_STR_ARG(what)); +} void ESPHomeOTAComponent::log_start_(const LogString *phase) { char peername[socket::SOCKADDR_STR_LEN]; @@ -793,6 +819,7 @@ void ESPHomeOTAComponent::cleanup_connection_() { this->handshake_buf_pos_ = 0; this->ota_state_ = OTAState::IDLE; this->ota_features_ = 0; + this->remote_closed_ = false; this->backend_ = nullptr; #ifdef USE_OTA_PASSWORD this->cleanup_auth_(); diff --git a/esphome/components/esphome/ota/ota_esphome.h b/esphome/components/esphome/ota/ota_esphome.h index 92ba094c8d..6f04b78da5 100644 --- a/esphome/components/esphome/ota/ota_esphome.h +++ b/esphome/components/esphome/ota/ota_esphome.h @@ -134,6 +134,7 @@ class ESPHomeOTAComponent final : public ota::OTAComponent { void server_failed_(const LogString *msg); void log_socket_error_(const LogString *msg); void log_read_error_(const LogString *what); + bool client_left_before_start_(); void log_start_(const LogString *phase); void log_remote_closed_(const LogString *during); void cleanup_connection_(); @@ -186,6 +187,7 @@ class ESPHomeOTAComponent final : public ota::OTAComponent { OTAState ota_state_{OTAState::IDLE}; uint8_t handshake_buf_pos_{0}; uint8_t ota_features_{0}; + bool remote_closed_{false}; // the peer hung up cleanly during a blocking read #ifdef USE_OTA_PASSWORD uint8_t auth_buf_pos_{0}; uint8_t auth_type_{0}; // Store auth type to know which hasher to use diff --git a/tests/integration/test_host_ota.py b/tests/integration/test_host_ota.py index 88eb0168e3..56a685eac3 100644 --- a/tests/integration/test_host_ota.py +++ b/tests/integration/test_host_ota.py @@ -167,6 +167,33 @@ class _Device: assert self.proc.returncode is None, "process died on rejected OTA" +def _handshake_then_close(port: int, noise_psk: str) -> None: + """Negotiate and complete the Noise handshake like a key probe, then + hang up without sending an OTA type.""" + with socket.create_connection((LOCALHOST, port), timeout=5.0) as sock: + espota2.send_check(sock, espota2.MAGIC_BYTES, "magic bytes") + _, version = espota2.receive_exactly(sock, 2, "version", espota2.RESPONSE_OK) + features_to_send = ( + espota2.CLIENT_FEATURE_SUPPORTS_COMPRESSION + | espota2.CLIENT_FEATURE_SUPPORTS_SHA256_AUTH + | espota2.CLIENT_FEATURE_SUPPORTS_EXTENDED_PROTOCOL + | espota2.CLIENT_FEATURE_SUPPORTS_NOISE + ) + espota2.send_check(sock, features_to_send, "features") + espota2.receive_exactly(sock, 1, "features", espota2.RESPONSE_FEATURE_FLAGS) + (features,) = espota2.receive_exactly(sock, 1, "feature flags", None) + assert features & espota2.SERVER_FEATURE_SUPPORTS_NOISE + prologue = ( + espota2.NOISE_PROLOGUE_INIT + + bytes(espota2.MAGIC_BYTES) + + bytes([espota2.RESPONSE_OK, version, features_to_send]) + + bytes([espota2.RESPONSE_FEATURE_FLAGS, features]) + ) + noise = espota2.NoiseSocketWrapper(sock, noise_psk, prologue) + noise.do_handshake() + espota2.receive_exactly(noise, 1, "auth", espota2.RESPONSE_AUTH_OK) + + async def _provision_key( dev: _Device, api_client_connected: APIClientConnectedFactory ) -> None: @@ -221,16 +248,24 @@ async def test_host_ota_encrypted( compile_esphome: CompileFunction, reserved_tcp_port: tuple[int, socket.socket], ) -> None: - """Encrypted self-OTA succeeds; a plaintext upload to the same device fails.""" + """A client that leaves right after the handshake, as a key probe does, + is a clean close, not an OTA error; a plaintext upload is refused; an + encrypted self-OTA succeeds.""" pytest.importorskip("aioesphomeapi.noise") dev = _Device( *await _build( yaml_config, write_yaml_config, compile_esphome, reserved_tcp_port ) ) - async with run_binary(dev.binary_path, line_callback=dev.on_log) as (proc, _lines): + async with run_binary(dev.binary_path, line_callback=dev.on_log) as (proc, lines): dev.proc = proc await _wait_for_port(LOCALHOST, dev.api_port, PORT_WAIT_TIMEOUT) + await asyncio.get_running_loop().run_in_executor( + None, _handshake_then_close, dev.ota_port, API_KEY + ) + # The error path logs its warning instead of this line, never after it + await _wait_for_line(lines, "Client left after the handshake") + assert not [line for line in lines if "[W][esphome.ota" in line] await dev.refused_ota( None, None, "plaintext upload to an encrypted device must fail" )