From 3ee8aaf77c94a643dfb87e900edc97101d1b81ec Mon Sep 17 00:00:00 2001 From: "J. Nick Koston" Date: Wed, 12 Aug 2026 18:54:57 -0500 Subject: [PATCH] Address review feedback on retry behavior and diagnostics --- esphome/espota2.py | 135 ++++++++++++--------- tests/unit_tests/test_espota2.py | 193 ++++++++++++++++++++++++------- 2 files changed, 237 insertions(+), 91 deletions(-) diff --git a/esphome/espota2.py b/esphome/espota2.py index dffa2b13f2..9565cc3d31 100644 --- a/esphome/espota2.py +++ b/esphome/espota2.py @@ -1,6 +1,7 @@ from __future__ import annotations from collections.abc import Callable +import contextlib import gzip import hashlib import io @@ -8,7 +9,6 @@ import logging from pathlib import Path import secrets import socket -import sys import time from typing import Any @@ -225,8 +225,9 @@ def receive_exactly( check_error(data, expect) except OTAError as err: sock.close() - # Preserve the subclass (OTANetworkError vs OTAError) so callers can - # distinguish retryable network failures from device-reported errors. + # type(err) preserves OTANetworkError vs OTAError so callers can tell + # retryable network failures from device-reported errors; subclasses + # must accept a single message argument raise type(err)(f"receiving {msg}: {err}") from err while len(data) < amount: @@ -463,21 +464,25 @@ def perform_ota( offset = 0 progress = ProgressBar("Uploading") - while True: - chunk = upload_contents[offset : offset + UPLOAD_BLOCK_SIZE] - if not chunk: - break - offset += len(chunk) + try: + while True: + chunk = upload_contents[offset : offset + UPLOAD_BLOCK_SIZE] + if not chunk: + break + offset += len(chunk) - try: - sock.sendall(chunk) - if version >= OTA_VERSION_2_0: - receive_exactly(sock, 1, "chunk result", RESPONSE_CHUNK_OK) - except OSError as err: - sys.stderr.write("\n") - raise OTANetworkError(f"sending data: {err}") from err + try: + sock.sendall(chunk) + if version >= OTA_VERSION_2_0: + receive_exactly(sock, 1, "chunk result", RESPONSE_CHUNK_OK) + except OSError as err: + raise OTANetworkError(f"sending data: {err}") from err - progress.update(offset / upload_size) + progress.update(offset / upload_size) + except OTAError: + # Terminate the progress bar line before the error is logged + progress.done() + raise progress.done() # Enable nodelay for last checks @@ -486,9 +491,26 @@ def perform_ota( _LOGGER.info("Upload took %.2f seconds, waiting for result...", duration) - receive_exactly(sock, 1, "update receive result", RESPONSE_RECEIVE_OK) - receive_exactly(sock, 1, "update end result", RESPONSE_UPDATE_END_OK) - send_check(sock, RESPONSE_OK, "end acknowledgement") + # Once the device has the complete image it commits the update and + # reboots on its own; the exact commit point is not observable from + # here, so treat everything past the data phase as non-retryable. A + # re-upload could flash a device that already updated successfully. + try: + receive_exactly(sock, 1, "update receive result", RESPONSE_RECEIVE_OK) + receive_exactly(sock, 1, "update end result", RESPONSE_UPDATE_END_OK) + except OTANetworkError as err: + raise OTAError( + f"{err} (the device may have already committed the update and " + f"be rebooting; check whether it comes back with the new " + f"firmware before uploading again)" + ) from err + + try: + send_check(sock, RESPONSE_OK, "end acknowledgement") + except OTANetworkError as err: + # The device treats a missing end acknowledgement as non-fatal and is + # already rebooting into the new firmware, so the update succeeded + _LOGGER.warning("Failed sending end acknowledgement: %s", err) _LOGGER.info("OTA successful") @@ -524,48 +546,57 @@ def run_ota_impl_( ) raise OTAError(err) from err - for attempt in range(1, MAX_UPLOAD_ATTEMPTS + 1): - if attempt > 1: + # Every address is tried at least once and the budget grants + # MAX_UPLOAD_ATTEMPTS - 1 extra retries, cycling through the addresses. + # Wait before an attempt when the previous one actually reached the + # device, or when revisiting an address, so a flaky link can recover and + # the device can clean up a half-open connection (its handshake watchdog + # runs at 20s); moving on to the next address family stays immediate. + total_attempts = len(res) + MAX_UPLOAD_ATTEMPTS - 1 + last_error = "" + reached_device = False + for attempt in range(total_attempts): + af, socktype, _, _, sa = res[attempt % len(res)] + if reached_device or attempt >= len(res): _LOGGER.info( "Retrying in %.0f seconds (attempt %d of %d)...", UPLOAD_RETRY_DELAY, - attempt, - MAX_UPLOAD_ATTEMPTS, + attempt + 1, + total_attempts, ) time.sleep(UPLOAD_RETRY_DELAY) - for af, socktype, _, _, sa in res: - _LOGGER.info("Connecting to %s port %s...", sa[0], sa[1]) - sock = socket.socket(af, socktype) - sock.settimeout(20.0) + reached_device = False + _LOGGER.info("Connecting to %s port %s...", sa[0], sa[1]) + sock = socket.socket(af, socktype) + sock.settimeout(20.0) + try: + sock.connect(sa) + except OSError as err: + sock.close() + _LOGGER.warning("Connecting to %s port %s failed: %s", sa[0], sa[1], err) + last_error = f"connecting to {sa[0]} failed: {err}" + continue + + _LOGGER.info("Connected to %s", sa[0]) + reached_device = True + with contextlib.closing(sock), Path(filename).open("rb") as file_handle: try: - sock.connect(sa) - except OSError as err: - sock.close() - _LOGGER.warning( - "Connecting to %s port %s failed: %s", sa[0], sa[1], err - ) + perform_ota(sock, password, file_handle, filename, ota_type) + except OTANetworkError as err: + # Transient network failure; retry + last_error = str(err) + _LOGGER.warning(last_error) continue + except OTAError as err: + # Device-reported error (wrong password, wrong flash size, ...); + # retrying cannot succeed, so fail immediately + _LOGGER.error(str(err)) + return 1, None - _LOGGER.info("Connected to %s", sa[0]) - with Path(filename).open("rb") as file_handle: - try: - perform_ota(sock, password, file_handle, filename, ota_type) - except OTANetworkError as err: - # Transient network failure; try the next address or attempt - _LOGGER.warning(str(err)) - continue - except OTAError as err: - # Device-reported error (wrong password, wrong flash size, - # ...); retrying cannot succeed, so fail immediately - _LOGGER.error(str(err)) - return 1, None - finally: - sock.close() + # Successfully uploaded to sa[0] + return 0, sa[0] - # Successfully uploaded to sa[0] - return 0, sa[0] - - _LOGGER.error("Connection failed after %d attempts.", MAX_UPLOAD_ATTEMPTS) + _LOGGER.error("Upload failed after %d attempts: %s", total_attempts, last_error) return 1, None diff --git a/tests/unit_tests/test_espota2.py b/tests/unit_tests/test_espota2.py index f78663ff35..b55758fa19 100644 --- a/tests/unit_tests/test_espota2.py +++ b/tests/unit_tests/test_espota2.py @@ -44,13 +44,17 @@ def mock_file() -> io.BytesIO: @pytest.fixture -def mock_time() -> Generator[None]: +def mock_sleep() -> Generator[Mock]: + """Mock time.sleep so delays don't slow down tests.""" + with patch("time.sleep") as mock: + yield mock + + +@pytest.fixture +def mock_time(mock_sleep: Mock) -> Generator[None]: """Mock time-related functions for consistent testing.""" # Provide enough values for multiple calls (tests may call perform_ota multiple times) - with ( - patch("time.sleep"), - patch("time.perf_counter", side_effect=[0, 1, 0, 1, 0, 1]), - ): + with patch("time.perf_counter", side_effect=[0, 1, 0, 1, 0, 1]): yield @@ -69,13 +73,6 @@ def mock_token_hex() -> Generator[Mock]: yield mock -@pytest.fixture -def mock_sleep() -> Generator[Mock]: - """Mock time.sleep so retry delays don't slow down tests.""" - with patch("esphome.espota2.time.sleep") as mock: - yield mock - - @pytest.fixture def mock_resolve_ip() -> Generator[Mock]: """Mock resolve_ip_address for testing.""" @@ -86,6 +83,28 @@ def mock_resolve_ip() -> Generator[Mock]: yield mock +DUAL_STACK_SA6 = ("2001:db8::1", 3232, 0, 0) +DUAL_STACK_SA4 = ("192.168.1.100", 3232) + + +@pytest.fixture +def mock_resolve_ip_dual(mock_resolve_ip: Mock) -> Mock: + """Make resolve_ip_address return an IPv6 and an IPv4 address.""" + mock_resolve_ip.return_value = [ + (socket.AF_INET6, socket.SOCK_STREAM, 0, "", DUAL_STACK_SA6), + (socket.AF_INET, socket.SOCK_STREAM, 0, "", DUAL_STACK_SA4), + ] + return mock_resolve_ip + + +@pytest.fixture +def firmware_file(tmp_path: Path) -> Path: + """Create a firmware file on disk for run_ota_impl_ tests.""" + firmware = tmp_path / "firmware.bin" + firmware.write_bytes(b"firmware content") + return firmware + + @pytest.fixture def mock_perform_ota() -> Generator[Mock]: """Mock perform_ota function for testing.""" @@ -560,17 +579,21 @@ def test_perform_ota_upload_error(mock_socket: Mock, mock_file: io.BytesIO) -> N @pytest.mark.usefixtures("mock_time") -def test_perform_ota_chunk_send_error(mock_socket: Mock, mock_file: io.BytesIO) -> None: - """Test OTA raises the retryable OTANetworkError when sending a chunk fails.""" - recv_responses = [ +def _no_auth_handshake(version: int) -> list[bytes]: + """Recv responses for a handshake without auth, up to the MD5 check.""" + return [ bytes([espota2.RESPONSE_OK]), # First byte of version response - bytes([espota2.OTA_VERSION_2_0]), # Version number + bytes([version]), # Version number bytes([espota2.RESPONSE_HEADER_OK]), # Features response bytes([espota2.RESPONSE_AUTH_OK]), # No auth required bytes([espota2.RESPONSE_UPDATE_PREPARE_OK]), # Binary size OK bytes([espota2.RESPONSE_BIN_MD5_OK]), # MD5 checksum OK ] - mock_socket.recv.side_effect = recv_responses + + +def test_perform_ota_chunk_send_error(mock_socket: Mock, mock_file: io.BytesIO) -> None: + """Test OTA raises the retryable OTANetworkError when sending a chunk fails.""" + mock_socket.recv.side_effect = _no_auth_handshake(espota2.OTA_VERSION_2_0) # Sends before the data phase: magic bytes, features, binary size, MD5; # fail on the fifth sendall, the first firmware chunk mock_socket.sendall.side_effect = [None] * 4 + [OSError("Broken pipe")] @@ -579,6 +602,44 @@ def test_perform_ota_chunk_send_error(mock_socket: Mock, mock_file: io.BytesIO) espota2.perform_ota(mock_socket, None, mock_file, "test.bin") +@pytest.mark.usefixtures("mock_time") +def test_perform_ota_post_commit_failure_not_retryable( + mock_socket: Mock, mock_file: io.BytesIO +) -> None: + """Test a network failure after the device committed is a plain OTAError.""" + mock_socket.recv.side_effect = [ + *_no_auth_handshake(espota2.OTA_VERSION_1_0), + bytes([espota2.RESPONSE_RECEIVE_OK]), # Device received everything + OSError("Connection reset"), # Connection lost waiting for end result + ] + + with pytest.raises(espota2.OTAError, match="receiving update end result") as exc: + espota2.perform_ota(mock_socket, None, mock_file, "test.bin") + + # Must not be the retryable kind; the device is already rebooting + assert not isinstance(exc.value, espota2.OTANetworkError) + + +@pytest.mark.usefixtures("mock_time") +def test_perform_ota_end_ack_send_failure_is_success( + mock_socket: Mock, mock_file: io.BytesIO +) -> None: + """Test a send failure on the final acknowledgement does not fail the OTA.""" + mock_socket.recv.side_effect = [ + *_no_auth_handshake(espota2.OTA_VERSION_1_0), + bytes([espota2.RESPONSE_RECEIVE_OK]), # Device received everything + bytes([espota2.RESPONSE_UPDATE_END_OK]), # Update committed + ] + # Sends: magic bytes, features, binary size, MD5, one firmware chunk; + # fail on the sixth sendall, the end acknowledgement + mock_socket.sendall.side_effect = [None] * 5 + [OSError("Broken pipe")] + + # Must not raise; the device treats a missing acknowledgement as non-fatal + espota2.perform_ota(mock_socket, None, mock_file, "test.bin") + + assert mock_socket.sendall.call_count == 6 + + @pytest.mark.usefixtures("mock_socket_constructor", "mock_resolve_ip") def test_run_ota_impl_successful( mock_socket: Mock, tmp_path: Path, mock_perform_ota: Mock @@ -614,22 +675,19 @@ def test_run_ota_impl_successful( @pytest.mark.usefixtures("mock_socket_constructor", "mock_resolve_ip") def test_run_ota_impl_connection_failed( - mock_socket: Mock, tmp_path: Path, mock_sleep: Mock + mock_socket: Mock, firmware_file: Path, mock_sleep: Mock ) -> None: """Test run_ota_impl_ retries when connection fails and eventually gives up.""" mock_socket.connect.side_effect = OSError("Connection refused") - # Create a real firmware file - firmware_file = tmp_path / "firmware.bin" - firmware_file.write_bytes(b"firmware content") - result_code, result_host = espota2.run_ota_impl_( "test.local", 3232, "password", str(firmware_file) ) assert result_code == 1 assert result_host is None - # One connect attempt per retry round, with a delay between rounds + # A single address gets the whole attempt budget, with a delay before + # each revisit assert mock_socket.connect.call_count == espota2.MAX_UPLOAD_ATTEMPTS assert mock_socket.close.call_count == espota2.MAX_UPLOAD_ATTEMPTS assert mock_sleep.call_count == espota2.MAX_UPLOAD_ATTEMPTS - 1 @@ -638,14 +696,11 @@ def test_run_ota_impl_connection_failed( @pytest.mark.usefixtures("mock_socket_constructor", "mock_resolve_ip") def test_run_ota_impl_connect_retry_succeeds( - mock_socket: Mock, tmp_path: Path, mock_perform_ota: Mock, mock_sleep: Mock + mock_socket: Mock, firmware_file: Path, mock_perform_ota: Mock, mock_sleep: Mock ) -> None: """Test run_ota_impl_ succeeds when a retry connects after a failed attempt.""" mock_socket.connect.side_effect = [OSError("Connection timed out"), None] - firmware_file = tmp_path / "firmware.bin" - firmware_file.write_bytes(b"firmware content") - result_code, result_host = espota2.run_ota_impl_( "test.local", 3232, "password", str(firmware_file) ) @@ -659,7 +714,7 @@ def test_run_ota_impl_connect_retry_succeeds( @pytest.mark.usefixtures("mock_socket_constructor", "mock_resolve_ip") def test_run_ota_impl_network_error_retry_succeeds( - mock_socket: Mock, tmp_path: Path, mock_perform_ota: Mock, mock_sleep: Mock + mock_socket: Mock, firmware_file: Path, mock_perform_ota: Mock, mock_sleep: Mock ) -> None: """Test run_ota_impl_ retries after a network error during the upload.""" mock_perform_ota.side_effect = [ @@ -667,9 +722,6 @@ def test_run_ota_impl_network_error_retry_succeeds( None, ] - firmware_file = tmp_path / "firmware.bin" - firmware_file.write_bytes(b"firmware content") - result_code, result_host = espota2.run_ota_impl_( "test.local", 3232, "password", str(firmware_file) ) @@ -682,14 +734,11 @@ def test_run_ota_impl_network_error_retry_succeeds( @pytest.mark.usefixtures("mock_socket_constructor", "mock_resolve_ip") def test_run_ota_impl_network_error_exhausts_attempts( - mock_socket: Mock, tmp_path: Path, mock_perform_ota: Mock, mock_sleep: Mock + mock_socket: Mock, firmware_file: Path, mock_perform_ota: Mock, mock_sleep: Mock ) -> None: """Test run_ota_impl_ gives up after all attempts hit network errors.""" mock_perform_ota.side_effect = espota2.OTANetworkError("sending data: broken pipe") - firmware_file = tmp_path / "firmware.bin" - firmware_file.write_bytes(b"firmware content") - result_code, result_host = espota2.run_ota_impl_( "test.local", 3232, "password", str(firmware_file) ) @@ -700,18 +749,84 @@ def test_run_ota_impl_network_error_exhausts_attempts( assert mock_sleep.call_count == espota2.MAX_UPLOAD_ATTEMPTS - 1 +@pytest.mark.usefixtures("mock_socket_constructor", "mock_resolve_ip_dual") +def test_run_ota_impl_multiple_addresses_cycle( + mock_socket: Mock, firmware_file: Path, mock_sleep: Mock +) -> None: + """Test run_ota_impl_ visits every address and cycles for the retries.""" + mock_socket.connect.side_effect = OSError("No route to host") + + result_code, result_host = espota2.run_ota_impl_( + "test.local", 3232, "password", str(firmware_file) + ) + + assert result_code == 1 + assert result_host is None + # Every address gets a first visit plus MAX_UPLOAD_ATTEMPTS - 1 retries + assert mock_socket.connect.call_args_list == [ + call(DUAL_STACK_SA6), + call(DUAL_STACK_SA4), + call(DUAL_STACK_SA6), + call(DUAL_STACK_SA4), + ] + # No connect ever reached the device, so the delay only applies before + # the revisits + assert mock_sleep.call_count == 2 + + +@pytest.mark.usefixtures("mock_socket_constructor", "mock_resolve_ip_dual") +def test_run_ota_impl_second_address_succeeds_without_delay( + mock_socket: Mock, + firmware_file: Path, + mock_perform_ota: Mock, + mock_sleep: Mock, +) -> None: + """Test run_ota_impl_ falls through to the next address with no pause.""" + mock_socket.connect.side_effect = [OSError("No route to host"), None] + + result_code, result_host = espota2.run_ota_impl_( + "test.local", 3232, "password", str(firmware_file) + ) + + assert result_code == 0 + assert result_host == "192.168.1.100" + mock_sleep.assert_not_called() + mock_perform_ota.assert_called_once() + + +@pytest.mark.usefixtures("mock_socket_constructor", "mock_resolve_ip_dual") +def test_run_ota_impl_pauses_after_reaching_device( + mock_socket: Mock, + firmware_file: Path, + mock_perform_ota: Mock, + mock_sleep: Mock, +) -> None: + """Test run_ota_impl_ pauses before the next address once the device was reached.""" + mock_perform_ota.side_effect = [ + espota2.OTANetworkError("sending data: connection reset"), + None, + ] + + result_code, result_host = espota2.run_ota_impl_( + "test.local", 3232, "password", str(firmware_file) + ) + + assert result_code == 0 + assert result_host == "192.168.1.100" + # The first attempt reached the device, so the next one waits first even + # though it targets a fresh address + mock_sleep.assert_called_once_with(espota2.UPLOAD_RETRY_DELAY) + + @pytest.mark.usefixtures("mock_socket_constructor", "mock_resolve_ip") def test_run_ota_impl_device_error_not_retried( - mock_socket: Mock, tmp_path: Path, mock_perform_ota: Mock, mock_sleep: Mock + mock_socket: Mock, firmware_file: Path, mock_perform_ota: Mock, mock_sleep: Mock ) -> None: """Test run_ota_impl_ fails immediately on a device-reported error.""" mock_perform_ota.side_effect = espota2.OTAError( "Authentication invalid. Is the password correct?" ) - firmware_file = tmp_path / "firmware.bin" - firmware_file.write_bytes(b"firmware content") - result_code, result_host = espota2.run_ota_impl_( "test.local", 3232, "password", str(firmware_file) )