diff --git a/esphome/espota2.py b/esphome/espota2.py index fa15c1dda2..479429fcc8 100644 --- a/esphome/espota2.py +++ b/esphome/espota2.py @@ -76,6 +76,12 @@ _SUPPORTED_OTA_TYPES: frozenset[int] = frozenset( UPLOAD_BLOCK_SIZE = 8192 UPLOAD_BUFFER_SIZE = UPLOAD_BLOCK_SIZE * 8 +# Flaky Wi-Fi links often drop the first OTA attempt, and the device may need time +# to clean up a half-open connection (its handshake watchdog runs at 20s) before it +# accepts a new one, so wait between attempts instead of failing the upload outright. +MAX_UPLOAD_ATTEMPTS = 3 +UPLOAD_RETRY_DELAY = 5.0 + _LOGGER = logging.getLogger(__name__) # Authentication method lookup table: response -> (hash_func, nonce_size, name) @@ -171,6 +177,10 @@ class OTAError(EsphomeError): pass +class OTANetworkError(OTAError): + """Network-level OTA failure (timeout, reset, closed connection); retrying may succeed.""" + + def recv_decode( sock: socket.socket, amount: int, decode: bool = True ) -> bytes | list[int]: @@ -209,19 +219,21 @@ def receive_exactly( try: data += recv_decode(sock, 1, decode=decode) # type: ignore[operator] except OSError as err: - raise OTAError(f"receiving {msg} response: {err}") from err + raise OTANetworkError(f"receiving {msg} response: {err}") from err try: check_error(data, expect) except OTAError as err: sock.close() - raise OTAError(f"receiving {msg}: {err}") from err + # Preserve the subclass (OTANetworkError vs OTAError) so callers can + # distinguish retryable network failures from device-reported errors. + raise type(err)(f"receiving {msg}: {err}") from err while len(data) < amount: try: data += recv_decode(sock, amount - len(data), decode=decode) # type: ignore[operator] except OSError as err: - raise OTAError(f"receiving {msg}: {err}") from err + raise OTANetworkError(f"receiving {msg}: {err}") from err return data @@ -237,7 +249,7 @@ def check_error(data: list[int] | bytes, expect: int | list[int] | None) -> None # accept-any-response reads (e.g. feature negotiation, auth nonces) would be # silently passed through and surface later as cryptic decode/timeout failures. if not data: - raise OTAError( + raise OTANetworkError( "Device closed connection without responding. " "This may indicate the device ran out of memory, " "a network issue, or the connection was interrupted." @@ -274,7 +286,7 @@ def send_check( sock.sendall(data) except OSError as err: - raise OTAError(f"sending {msg}: {err}") from err + raise OTANetworkError(f"sending {msg}: {err}") from err def perform_ota( @@ -461,7 +473,7 @@ def perform_ota( receive_exactly(sock, 1, "chunk result", RESPONSE_CHUNK_OK) except OSError as err: sys.stderr.write("\n") - raise OTAError(f"sending data: {err}") from err + raise OTANetworkError(f"sending data: {err}") from err progress.update(offset / upload_size) progress.done() @@ -510,32 +522,49 @@ def run_ota_impl_( ) raise OTAError(err) from err - for r in res: - af, socktype, _, _, sa = r - _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.error("Connecting to %s port %s failed: %s", sa[0], sa[1], err) - continue - - _LOGGER.info("Connected to %s", sa[0]) - with Path(filename).open("rb") as file_handle: + for attempt in range(1, MAX_UPLOAD_ATTEMPTS + 1): + if attempt > 1: + _LOGGER.info( + "Retrying in %.0f seconds (attempt %d of %d)...", + UPLOAD_RETRY_DELAY, + attempt, + MAX_UPLOAD_ATTEMPTS, + ) + time.sleep(UPLOAD_RETRY_DELAY) + for r in res: + af, socktype, _, _, sa = r + _LOGGER.info("Connecting to %s port %s...", sa[0], sa[1]) + sock = socket.socket(af, socktype) + sock.settimeout(20.0) try: - perform_ota(sock, password, file_handle, filename, ota_type) - except OTAError as err: - _LOGGER.error(str(err)) - return 1, None - finally: + sock.connect(sa) + except OSError as err: sock.close() + _LOGGER.warning( + "Connecting to %s port %s failed: %s", sa[0], sa[1], err + ) + continue - # Successfully uploaded to sa[0] - return 0, sa[0] + _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() - _LOGGER.error("Connection failed.") + # Successfully uploaded to sa[0] + return 0, sa[0] + + _LOGGER.error("Connection failed after %d attempts.", MAX_UPLOAD_ATTEMPTS) return 1, None diff --git a/tests/unit_tests/test_espota2.py b/tests/unit_tests/test_espota2.py index 9413fbcf29..3bf65f62a0 100644 --- a/tests/unit_tests/test_espota2.py +++ b/tests/unit_tests/test_espota2.py @@ -69,6 +69,13 @@ 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.""" @@ -147,10 +154,32 @@ def test_receive_exactly_socket_error(mock_socket: Mock) -> None: """Test receive_exactly handles socket errors.""" mock_socket.recv.side_effect = OSError("Connection reset") - with pytest.raises(espota2.OTAError, match="receiving test response"): + with pytest.raises(espota2.OTANetworkError, match="receiving test response"): espota2.receive_exactly(mock_socket, 1, "test", espota2.RESPONSE_OK) +def test_receive_exactly_closed_connection_is_network_error(mock_socket: Mock) -> None: + """Test receive_exactly raises OTANetworkError when the device closes the connection.""" + mock_socket.recv.return_value = b"" + + with pytest.raises( + espota2.OTANetworkError, match="Device closed connection without responding" + ): + espota2.receive_exactly(mock_socket, 1, "test", espota2.RESPONSE_OK) + + mock_socket.close.assert_called_once() + + +def test_receive_exactly_device_error_is_not_network_error(mock_socket: Mock) -> None: + """Test receive_exactly keeps device-reported errors as plain OTAError.""" + mock_socket.recv.return_value = bytes([espota2.RESPONSE_ERROR_AUTH_INVALID]) + + with pytest.raises(espota2.OTAError) as exc_info: + espota2.receive_exactly(mock_socket, 1, "auth", [espota2.RESPONSE_OK]) + + assert not isinstance(exc_info.value, espota2.OTANetworkError) + + @pytest.mark.parametrize( ("error_code", "expected_msg"), [ @@ -226,6 +255,12 @@ def test_check_error_unexpected_response() -> None: espota2.check_error([0x7F], [espota2.RESPONSE_OK, espota2.RESPONSE_AUTH_OK]) +def test_check_error_empty_data_is_network_error() -> None: + """Test check_error raises the retryable OTANetworkError subclass on empty data.""" + with pytest.raises(espota2.OTANetworkError): + espota2.check_error(b"", espota2.RESPONSE_OK) + + def test_check_error_empty_data() -> None: """Test check_error raises error when device closes connection without responding.""" with pytest.raises( @@ -564,8 +599,10 @@ 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) -> None: - """Test run_ota_impl_ when connection fails.""" +def test_run_ota_impl_connection_failed( + mock_socket: Mock, tmp_path: 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 @@ -578,7 +615,97 @@ def test_run_ota_impl_connection_failed(mock_socket: Mock, tmp_path: Path) -> No assert result_code == 1 assert result_host is None - mock_socket.close.assert_called_once() + # One connect attempt per retry round, with a delay between rounds + 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 + mock_sleep.assert_called_with(espota2.UPLOAD_RETRY_DELAY) + + +@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 +) -> 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) + ) + + assert result_code == 0 + assert result_host == "192.168.1.100" + assert mock_socket.connect.call_count == 2 + mock_sleep.assert_called_once_with(espota2.UPLOAD_RETRY_DELAY) + mock_perform_ota.assert_called_once() + + +@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 +) -> None: + """Test run_ota_impl_ retries after a network error during the upload.""" + mock_perform_ota.side_effect = [ + espota2.OTANetworkError("receiving features: Device closed connection"), + 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) + ) + + assert result_code == 0 + assert result_host == "192.168.1.100" + assert mock_perform_ota.call_count == 2 + mock_sleep.assert_called_once_with(espota2.UPLOAD_RETRY_DELAY) + + +@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 +) -> 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) + ) + + assert result_code == 1 + assert result_host is None + assert mock_perform_ota.call_count == espota2.MAX_UPLOAD_ATTEMPTS + assert mock_sleep.call_count == espota2.MAX_UPLOAD_ATTEMPTS - 1 + + +@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 +) -> 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) + ) + + assert result_code == 1 + assert result_host is None + mock_perform_ota.assert_called_once() + mock_sleep.assert_not_called() def test_run_ota_impl_resolve_failed(tmp_path: Path, mock_resolve_ip: Mock) -> None: