[ota] Offer encryption with the api key so enabling it works over OTA (#18979)

This commit is contained in:
J. Nick Koston
2026-09-07 10:46:21 +12:00
committed by Jesse Hills
parent 18220e0b39
commit 95ab3fb4f2
44 changed files with 1342 additions and 443 deletions
+125 -11
View File
@@ -10,12 +10,15 @@ when the installed aioesphomeapi predates the noise module.
from __future__ import annotations
import base64
from collections.abc import Callable
import hashlib
import io
import logging
from pathlib import Path
import socket
import sys
import threading
from typing import Any
from unittest.mock import Mock, patch
import pytest
@@ -65,8 +68,12 @@ class FakeEncryptedDevice(threading.Thread):
offer_noise: bool = True,
require_noise: bool = True,
prologue_features_override: int | None = None,
connections: int = 1,
drop_handshakes: int = 0,
) -> None:
super().__init__(daemon=True)
self.connections = connections
self.drop_handshakes = drop_handshakes # hang up mid-handshake this many times
self.psk = psk
self.version = version
self.offer_noise = offer_noise
@@ -81,10 +88,11 @@ class FakeEncryptedDevice(threading.Thread):
def run(self) -> None:
try:
sock, _ = self.listener.accept()
sock.settimeout(10)
with sock:
self._serve(sock)
for _ in range(self.connections):
sock, _ = self.listener.accept()
sock.settimeout(10)
with sock:
self._serve(sock)
except Exception as err: # noqa: BLE001 - surfaced via join_and_check
self.error = err
finally:
@@ -109,8 +117,23 @@ class FakeEncryptedDevice(threading.Thread):
return
server_flags = espota2.SERVER_FEATURE_SUPPORTS_NOISE if self.offer_noise else 0
sock.sendall(bytes([espota2.RESPONSE_FEATURE_FLAGS, server_flags]))
if not (self.offer_noise and noise_negotiated):
return # the client fails closed; nothing further arrives
if not (noise_negotiated and self.offer_noise):
# A device that does not require encryption continues in
# plaintext whatever the client asked for, like older firmware
try:
self._transfer(
lambda byte: sock.sendall(bytes([byte])),
lambda length: _recv_exact(sock, length),
lambda remaining: _recv_exact(
sock, min(remaining, espota2.UPLOAD_BLOCK_SIZE)
),
)
except ConnectionError:
# A keyed client without fallback fails closed and hangs up
if noise_negotiated and not self.offer_noise:
return
raise
return
from cryptography.exceptions import InvalidTag
from noise.connection import NoiseConnection
@@ -134,6 +157,9 @@ class FakeEncryptedDevice(threading.Thread):
msg1 = _recv_frame(sock)
assert msg1[0] == 0x00
if self.drop_handshakes > 0:
self.drop_handshakes -= 1
return # a transport fault: the socket closes with no reply
try:
proto.read_message(msg1[1:])
except InvalidTag:
@@ -149,6 +175,20 @@ class FakeEncryptedDevice(threading.Thread):
assert len(plaintext) == length, "control units must be one per frame"
return plaintext
def recv_data(_remaining: int) -> bytes:
plaintext = proto.decrypt(_recv_frame(sock))
assert 0 < len(plaintext) <= espota2.NOISE_MAX_PLAINTEXT
return plaintext
self._transfer(send_byte, recv_unit, recv_data)
def _transfer(
self,
send_byte: Callable[[int], None],
recv_unit: Callable[[int], bytes],
recv_data: Callable[[int], bytes],
) -> None:
"""The post-handshake exchange, identical over both transports."""
send_byte(espota2.RESPONSE_AUTH_OK)
recv_unit(1) # ota type
size = int.from_bytes(recv_unit(4), "big")
@@ -159,9 +199,7 @@ class FakeEncryptedDevice(threading.Thread):
received = b""
acked = 0
while len(received) < size:
plaintext = proto.decrypt(_recv_frame(sock))
assert 0 < len(plaintext) <= espota2.NOISE_MAX_PLAINTEXT
received += plaintext
received += recv_data(size - len(received))
if self.version >= espota2.OTA_VERSION_2_0:
while acked + espota2.UPLOAD_BLOCK_SIZE <= len(received) or (
len(received) == size and acked < size
@@ -176,7 +214,10 @@ class FakeEncryptedDevice(threading.Thread):
def _upload(
device: FakeEncryptedDevice, firmware: bytes, noise_psk: str | None
device: FakeEncryptedDevice,
firmware: bytes,
noise_psk: str | None,
plaintext_fallback: bool = False,
) -> None:
device.start()
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
@@ -184,12 +225,35 @@ def _upload(
sock.connect(("127.0.0.1", device.port))
try:
espota2.perform_ota(
sock, None, io.BytesIO(firmware), Path("firmware.bin"), noise_psk=noise_psk
sock,
None,
io.BytesIO(firmware),
Path("firmware.bin"),
noise_psk=noise_psk,
plaintext_fallback=plaintext_fallback,
)
finally:
sock.close()
def _run_ota(
device: FakeEncryptedDevice, firmware: bytes, tmp_path: Path, noise_psk: str
) -> int:
"""Drive the retry loop, which is where the plaintext fallback reconnects."""
path = tmp_path / "firmware.bin"
path.write_bytes(firmware)
device.start()
rc, _ = espota2.run_ota(
"127.0.0.1",
device.port,
None,
path,
noise_psk=noise_psk,
plaintext_fallback=True,
)
return rc
def test_encrypted_upload_success() -> None:
"""A full encrypted v2 upload spanning several 8192-byte blocks."""
pytest.importorskip("aioesphomeapi.noise")
@@ -240,6 +304,56 @@ def test_client_fails_closed_when_device_lacks_encryption() -> None:
device.join_and_check()
# Remove before 2027.3.0
def test_fallback_when_device_does_not_offer(caplog: pytest.LogCaptureFixture) -> None:
"""The api key is tried opportunistically; an older device that cannot
encrypt still gets its update, with a warning."""
firmware = b"firmware"
device = FakeEncryptedDevice(offer_noise=False, require_noise=False)
with patch("time.sleep"), caplog.at_level(logging.WARNING):
_upload(device, firmware, PSK, plaintext_fallback=True)
device.join_and_check()
assert device.received == firmware
assert any("fallback is removed in 2027.3.0" in r.message for r in caplog.records)
# Remove before 2027.3.0
@pytest.mark.parametrize(
("device_kwargs", "expected_rc", "fell_back"),
[
# A wrong key against an offering device reconnects in plaintext
({"psk": OTHER_PSK, "require_noise": False, "connections": 2}, 0, True),
# The plaintext retry is refused by a device that requires encryption
({"psk": OTHER_PSK, "require_noise": True, "connections": 2}, 1, True),
# A dropped connection inside the handshake is retried encrypted
({"require_noise": False, "connections": 2, "drop_handshakes": 1}, 0, False),
# A second transport fault inside the handshake falls back
({"require_noise": False, "connections": 3, "drop_handshakes": 2}, 0, True),
],
ids=["wrong_key", "wrong_key_required", "one_fault", "two_faults"],
)
def test_fallback_through_the_retry_loop(
caplog: pytest.LogCaptureFixture,
tmp_path: Path,
device_kwargs: dict[str, Any],
expected_rc: int,
fell_back: bool,
) -> None:
pytest.importorskip("aioesphomeapi.noise")
firmware = b"firmware"
device = FakeEncryptedDevice(**device_kwargs)
with patch("time.sleep"), caplog.at_level(logging.WARNING):
rc = _run_ota(device, firmware, tmp_path, PSK)
device.join_and_check()
assert rc == expected_rc
assert (device.received == firmware) is (expected_rc == 0)
assert (
any("Retrying in plaintext" in r.message for r in caplog.records) is fell_back
)
if expected_rc == 1:
assert any("requires an encrypted OTA" in r.message for r in caplog.records)
def test_plaintext_client_gets_encryption_required_error() -> None:
"""A client without a key gets the device's 0x94 error message."""
device = FakeEncryptedDevice()
+108 -6
View File
@@ -2108,7 +2108,13 @@ def test_upload_program_ota_success(
tmp_path / ".esphome" / "build" / "test" / ".pioenvs" / "test" / "firmware.bin"
)
mock_run_ota.assert_called_once_with(
["192.168.1.100"], 3232, "secret", expected_firmware, OTA_TYPE_UPDATE_APP, None
["192.168.1.100"],
3232,
"secret",
expected_firmware,
OTA_TYPE_UPDATE_APP,
None,
plaintext_fallback=False,
)
@@ -2140,10 +2146,77 @@ def test_upload_program_ota_encryption_key(
tmp_path / ".esphome" / "build" / "test" / ".pioenvs" / "test" / "firmware.bin"
)
mock_run_ota.assert_called_once_with(
["192.168.1.100"], 3232, None, expected_firmware, OTA_TYPE_UPDATE_APP, key
["192.168.1.100"],
3232,
None,
expected_firmware,
OTA_TYPE_UPDATE_APP,
key,
plaintext_fallback=False,
)
def test_upload_program_ota_api_key_opportunistic(
mock_run_ota: Mock,
mock_get_port_type: Mock,
tmp_path: Path,
) -> None:
"""Without an ota encryption block the api key is tried with a plaintext
fallback (removed in 2027.3.0)."""
setup_core(platform=PLATFORM_ESP32, tmp_path=tmp_path)
mock_get_port_type.return_value = "NETWORK"
mock_run_ota.return_value = (0, "192.168.1.100")
key = "AAECAwQFBgcICQoLDA0ODxAREhMUFRYXGBkaGxwdHh8="
config = {
CONF_API: {CONF_ENCRYPTION: {CONF_KEY: key}},
CONF_OTA: [{CONF_PLATFORM: CONF_ESPHOME, CONF_PORT: 3232}],
}
exit_code, _ = upload_program(config, MockArgs(), ["192.168.1.100"])
assert exit_code == 0
expected_firmware = (
tmp_path / ".esphome" / "build" / "test" / ".pioenvs" / "test" / "firmware.bin"
)
mock_run_ota.assert_called_once_with(
["192.168.1.100"],
3232,
None,
expected_firmware,
OTA_TYPE_UPDATE_APP,
key,
plaintext_fallback=True,
)
@pytest.mark.parametrize(
"api_conf",
[{}, {CONF_ENCRYPTION: {}}],
ids=["no_encryption", "runtime_key"],
)
def test_upload_program_ota_no_usable_api_key_stays_plaintext(
mock_run_ota: Mock,
mock_get_port_type: Mock,
tmp_path: Path,
api_conf: dict[str, Any],
) -> None:
"""A missing or runtime provisioned api key gives the uploader nothing
to try."""
setup_core(platform=PLATFORM_ESP32, tmp_path=tmp_path)
mock_get_port_type.return_value = "NETWORK"
mock_run_ota.return_value = (0, "192.168.1.100")
config = {
CONF_API: api_conf,
CONF_OTA: [{CONF_PLATFORM: CONF_ESPHOME, CONF_PORT: 3232}],
}
exit_code, _ = upload_program(config, MockArgs(), ["192.168.1.100"])
assert exit_code == 0
assert mock_run_ota.call_args.args[5] is None
assert mock_run_ota.call_args.kwargs == {"plaintext_fallback": False}
def test_upload_program_ota_encryption_without_key_fails_closed(
mock_run_ota: Mock,
mock_get_port_type: Mock,
@@ -2194,7 +2267,13 @@ def test_upload_program_ota_with_file_arg(
assert exit_code == 0
assert host == "192.168.1.100"
mock_run_ota.assert_called_once_with(
["192.168.1.100"], 3232, None, Path("custom.bin"), OTA_TYPE_UPDATE_APP, None
["192.168.1.100"],
3232,
None,
Path("custom.bin"),
OTA_TYPE_UPDATE_APP,
None,
plaintext_fallback=False,
)
@@ -2250,6 +2329,7 @@ def test_upload_program_ota_partition_table_with_file_arg(
partition_file,
OTA_TYPE_UPDATE_PARTITION_TABLE,
None,
plaintext_fallback=False,
)
@@ -2312,6 +2392,7 @@ def test_upload_program_ota_partition_table_mqttip(
partition_file,
OTA_TYPE_UPDATE_PARTITION_TABLE,
None,
plaintext_fallback=False,
)
@@ -2500,6 +2581,7 @@ def test_upload_program_ota_bootloader_with_file_arg(
bootloader_file,
OTA_TYPE_UPDATE_BOOTLOADER,
None,
plaintext_fallback=False,
)
@@ -2988,7 +3070,13 @@ def test_upload_program_ota_with_mqtt_resolution(
tmp_path / ".esphome" / "build" / "test" / ".pioenvs" / "test" / "firmware.bin"
)
mock_run_ota.assert_called_once_with(
["192.168.1.100"], 3232, None, expected_firmware, OTA_TYPE_UPDATE_APP, None
["192.168.1.100"],
3232,
None,
expected_firmware,
OTA_TYPE_UPDATE_APP,
None,
plaintext_fallback=False,
)
@@ -3038,7 +3126,13 @@ def test_upload_program_ota_with_mqtt_empty_broker(
tmp_path / ".esphome" / "build" / "test" / ".pioenvs" / "test" / "firmware.bin"
)
mock_run_ota.assert_called_once_with(
["192.168.1.50"], 3232, None, expected_firmware, OTA_TYPE_UPDATE_APP, None
["192.168.1.50"],
3232,
None,
expected_firmware,
OTA_TYPE_UPDATE_APP,
None,
plaintext_fallback=False,
)
# Verify warning was logged
assert "MQTT IP discovery failed" in caplog.text
@@ -5211,6 +5305,7 @@ def test_upload_program_ota_static_ip_with_mqttip(
expected_firmware,
OTA_TYPE_UPDATE_APP,
None,
plaintext_fallback=False,
)
@@ -5261,6 +5356,7 @@ def test_upload_program_ota_multiple_mqttip_resolves_once(
expected_firmware,
OTA_TYPE_UPDATE_APP,
None,
plaintext_fallback=False,
)
@@ -5438,7 +5534,13 @@ def test_upload_program_ota_mqtt_timeout_fallback(
tmp_path / ".esphome" / "build" / "test" / ".pioenvs" / "test" / "firmware.bin"
)
mock_run_ota.assert_called_once_with(
["192.168.1.100"], 3232, None, expected_firmware, OTA_TYPE_UPDATE_APP, None
["192.168.1.100"],
3232,
None,
expected_firmware,
OTA_TYPE_UPDATE_APP,
None,
plaintext_fallback=False,
)
+25 -6
View File
@@ -37,7 +37,6 @@ def wizard_answers() -> list[str]:
"nodemcuv2", # board
"SSID", # ssid
"psk", # wifi password
"", # ota password (empty for no password)
]
@@ -101,6 +100,25 @@ def test_config_file_should_include_ota(default_config: dict[str, Any]):
assert "ota:" in config
def test_config_file_should_use_encryption_when_api_key_set(
default_config: dict[str, Any],
):
"""
With an API encryption key and no OTA password the OTA block reuses the key
"""
# Given
default_config["api_encryption_key"] = (
"AAECAwQFBgcICQoLDA0ODxAREhMUFRYXGBkaGxwdHh8="
)
# When
config = wz.wizard_file(**default_config)
# Then
assert "ota:\n - platform: esphome\n encryption:" in config
assert "password" not in config.split("ota:")[1].split("wifi:")[0]
def test_config_file_should_include_ota_when_password_set(
default_config: dict[str, Any],
):
@@ -630,15 +648,15 @@ def test_wizard_write_protects_existing_config(
assert config_file.read_text() == original_content
def test_wizard_accepts_ota_password(
def test_wizard_uses_the_api_key_for_ota(
tmp_path: Path, monkeypatch: MonkeyPatch, wizard_answers: list[str]
):
"""
The wizard should pass ota_password to wizard_write when the user provides one
The wizard generates an api key and does not ask for an OTA password;
the key secures OTA updates
"""
# Given
wizard_answers[5] = "my_ota_password" # Set OTA password
config_file = tmp_path / "test.yaml"
input_mock = MagicMock(side_effect=wizard_answers)
monkeypatch.setattr("builtins.input", input_mock)
@@ -653,8 +671,9 @@ def test_wizard_accepts_ota_password(
# Then
assert retval == 0
call_kwargs = wizard_write_mock.call_args.kwargs
assert "ota_password" in call_kwargs
assert call_kwargs["ota_password"] == "my_ota_password"
assert "api_encryption_key" in call_kwargs
assert "ota_password" not in call_kwargs
assert input_mock.call_count == len(wizard_answers)
def test_wizard_accepts_rpipico_board(tmp_path: Path, monkeypatch: MonkeyPatch):