[ota] Retry uploads that fail from network errors

This commit is contained in:
J. Nick Koston
2026-08-12 18:23:00 -05:00
parent 99677390e0
commit 7b0541cd23
2 changed files with 188 additions and 32 deletions
+57 -28
View File
@@ -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
+131 -4
View File
@@ -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: