[ota] Offer encryption with the api key so enabling it works over OTA

This commit is contained in:
J. Nick Koston
2026-09-05 11:13:43 +02:00
parent d1829c495d
commit 029f6d4bc3
17 changed files with 425 additions and 32 deletions
@@ -2,6 +2,7 @@
from __future__ import annotations
from collections.abc import Callable
import logging
from typing import Any
@@ -11,6 +12,7 @@ from esphome import config_validation as cv
from esphome.components.esphome.ota import (
AUTO_LOAD,
FILTER_SOURCE_FILES,
_api_static_key,
_validate_no_password_with_encryption,
ota_esphome_final_validate,
)
@@ -386,6 +388,84 @@ def test_filter_source_files_excludes_noise_without_encryption() -> None:
CORE.config = old_config
def test_filter_source_files_keeps_noise_for_static_api_key() -> None:
"""A static api key makes the device offer encryption, so the transport
compiles even without an ota encryption block."""
old_config = CORE.config
ota = [_make_ota_config(port=3232)]
try:
CORE.config = {CONF_API: {CONF_ENCRYPTION: {CONF_KEY: API_KEY}}, CONF_OTA: ota}
assert FILTER_SOURCE_FILES() == []
# A runtime provisioned or all-zeros api key has nothing to offer
CORE.config = {CONF_API: {CONF_ENCRYPTION: {}}, CONF_OTA: ota}
assert FILTER_SOURCE_FILES() == ["ota_esphome_noise.cpp"]
CORE.config = {
CONF_API: {CONF_ENCRYPTION: {CONF_KEY: ZEROS_KEY}},
CONF_OTA: ota,
}
assert FILTER_SOURCE_FILES() == ["ota_esphome_noise.cpp"]
CORE.config = {CONF_API: {}, CONF_OTA: ota}
assert FILTER_SOURCE_FILES() == ["ota_esphome_noise.cpp"]
finally:
CORE.config = old_config
def test_api_static_key() -> None:
"""Only a real build-time api key can seed the encryption offer."""
assert _api_static_key({}) is None
assert _api_static_key({CONF_ENCRYPTION: {}}) is None
assert _api_static_key({CONF_ENCRYPTION: {CONF_KEY: ZEROS_KEY}}) is None
assert _api_static_key({CONF_ENCRYPTION: {CONF_KEY: API_KEY}}) == API_KEY
def _defines() -> set[str]:
return {define.name for define in CORE.defines}
def test_api_key_offers_encryption_without_requiring_it(
generate_main: Callable[[str], str],
) -> None:
"""An api key alone compiles the transport in and sets the psk, but the
device keeps accepting plaintext uploads."""
main_cpp = generate_main(
"tests/component_tests/ota/test_esphome_ota_api_key_offer.yaml"
)
assert "USE_OTA_ENCRYPTION" in _defines()
assert "USE_OTA_ENCRYPTION_REQUIRED" not in _defines()
assert "set_noise_psk(" in main_cpp
def test_api_key_offer_keeps_password(generate_main: Callable[[str], str]) -> None:
"""A password still guards plaintext uploads on an offering device."""
main_cpp = generate_main(
"tests/component_tests/ota/test_esphome_ota_api_key_offer_password.yaml"
)
assert {"USE_OTA_ENCRYPTION", "USE_OTA_PASSWORD"} <= _defines()
assert "USE_OTA_ENCRYPTION_REQUIRED" not in _defines()
assert "set_noise_psk(" in main_cpp
assert "set_auth_password(" in main_cpp
def test_encryption_block_requires_encryption(
generate_main: Callable[[str], str],
) -> None:
"""The ota encryption block is what makes the device refuse plaintext."""
main_cpp = generate_main(
"tests/component_tests/ota/test_esphome_ota_encryption_required.yaml"
)
assert {"USE_OTA_ENCRYPTION", "USE_OTA_ENCRYPTION_REQUIRED"} <= _defines()
assert "set_noise_psk(" in main_cpp
def test_runtime_api_key_offers_nothing(generate_main: Callable[[str], str]) -> None:
"""A key provisioned at runtime is unknown at build time, so no offer."""
main_cpp = generate_main(
"tests/component_tests/ota/test_esphome_ota_runtime_api_key.yaml"
)
assert "USE_OTA_ENCRYPTION" not in _defines()
assert "set_noise_psk(" not in main_cpp
def test_password_with_encryption_rejected() -> None:
"""The password and encryption options are mutually exclusive."""
config = {CONF_PASSWORD: "pw", CONF_ENCRYPTION: {CONF_KEY: API_KEY}}
@@ -0,0 +1,11 @@
esphome:
name: ota-offer
host:
api:
encryption:
key: "AAECAwQFBgcICQoLDA0ODxAREhMUFRYXGBkaGxwdHh8="
ota:
- platform: esphome
@@ -0,0 +1,12 @@
esphome:
name: ota-offer-password
host:
api:
encryption:
key: "AAECAwQFBgcICQoLDA0ODxAREhMUFRYXGBkaGxwdHh8="
ota:
- platform: esphome
password: "superlongpasswordthatnoonewillknow"
@@ -0,0 +1,12 @@
esphome:
name: ota-encryption-required
host:
api:
encryption:
key: "AAECAwQFBgcICQoLDA0ODxAREhMUFRYXGBkaGxwdHh8="
ota:
- platform: esphome
encryption:
@@ -0,0 +1,10 @@
esphome:
name: ota-runtime-key
host:
api:
encryption:
ota:
- platform: esphome
+12
View File
@@ -0,0 +1,12 @@
wifi:
ssid: MySSID
password: password1
api:
encryption:
key: "AAECAwQFBgcICQoLDA0ODxAREhMUFRYXGBkaGxwdHh8="
ota:
- platform: esphome
port: 3290
password: "superlongpasswordthatnoonewillknow"
@@ -0,0 +1,2 @@
packages:
ota: !include api_key_offer.yaml
@@ -0,0 +1,2 @@
packages:
ota: !include api_key_offer.yaml
@@ -0,0 +1,12 @@
esphome:
name: host-ota-test
host:
api:
encryption:
key: "AAECAwQFBgcICQoLDA0ODxAREhMUFRYXGBkaGxwdHh8="
ota:
- platform: esphome
port: __OTA_PORT__
password: "hunter2"
logger:
level: DEBUG
@@ -0,0 +1,11 @@
esphome:
name: host-ota-test
host:
api:
encryption:
key: "AAECAwQFBgcICQoLDA0ODxAREhMUFRYXGBkaGxwdHh8="
ota:
- platform: esphome
port: __OTA_PORT__
logger:
level: DEBUG
+137
View File
@@ -168,6 +168,143 @@ async def test_host_ota_encrypted(
assert proc.pid == pid_before
class _RebootCounter:
"""Counts safe reboots so a test can wait for the nth one."""
def __init__(self) -> None:
self._seen = asyncio.Event()
self.count = 0
def on_log(self, line: str) -> None:
if "Rebooting safely" in line:
self.count += 1
self._seen.set()
async def wait(self, count: int, timeout: float = 10.0) -> None:
async with asyncio.timeout(timeout):
while self.count < count:
self._seen.clear()
await self._seen.wait()
@pytest.mark.asyncio
async def test_host_ota_api_key_offers_encryption(
yaml_config: str,
write_yaml_config: ConfigWriter,
compile_esphome: CompileFunction,
reserved_tcp_port: tuple[int, socket.socket],
) -> None:
"""With only an api key the device takes both a plaintext upload and an
encrypted one using that key, which is the enablement path for
`ota: encryption:`."""
pytest.importorskip("aioesphomeapi.noise")
api_key = "AAECAwQFBgcICQoLDA0ODxAREhMUFRYXGBkaGxwdHh8="
api_port, api_socket = reserved_tcp_port
with _reserve_port() as (ota_port, ota_socket):
yaml_config = yaml_config.replace("__OTA_PORT__", str(ota_port))
config_path = await write_yaml_config(yaml_config)
binary_path = await compile_esphome(config_path)
api_socket.close()
ota_socket.close()
loop = asyncio.get_running_loop()
reboots = _RebootCounter()
async with run_binary(binary_path, line_callback=reboots.on_log) as (
proc,
lines,
):
await _wait_for_port(LOCALHOST, api_port, PORT_WAIT_TIMEOUT)
pid_before = proc.pid
rc, _ = await loop.run_in_executor(
None, espota2.run_ota, LOCALHOST, ota_port, None, binary_path
)
assert rc == 0, "plaintext upload to an offering device must succeed"
await reboots.wait(1)
await _wait_for_port(LOCALHOST, api_port, PORT_WAIT_TIMEOUT)
assert proc.pid == pid_before
rc, _ = await loop.run_in_executor(
None,
functools.partial(
espota2.run_ota,
LOCALHOST,
ota_port,
None,
binary_path,
noise_psk=api_key,
),
)
assert rc == 0, "encrypted upload with the api key must succeed"
await reboots.wait(2)
await _wait_for_port(LOCALHOST, api_port, PORT_WAIT_TIMEOUT)
assert proc.returncode is None, "process exited instead of execing"
assert proc.pid == pid_before
assert any("Encryption: offered" in line for line in lines)
@pytest.mark.asyncio
async def test_host_ota_api_key_offer_with_password(
yaml_config: str,
write_yaml_config: ConfigWriter,
compile_esphome: CompileFunction,
reserved_tcp_port: tuple[int, socket.socket],
) -> None:
"""The OTA password still guards plaintext uploads on an offering device
while the api key alone authenticates an encrypted one."""
pytest.importorskip("aioesphomeapi.noise")
api_key = "AAECAwQFBgcICQoLDA0ODxAREhMUFRYXGBkaGxwdHh8="
api_port, api_socket = reserved_tcp_port
with _reserve_port() as (ota_port, ota_socket):
yaml_config = yaml_config.replace("__OTA_PORT__", str(ota_port))
config_path = await write_yaml_config(yaml_config)
binary_path = await compile_esphome(config_path)
api_socket.close()
ota_socket.close()
loop = asyncio.get_running_loop()
reboots = _RebootCounter()
async with run_binary(binary_path, line_callback=reboots.on_log) as (
proc,
_lines,
):
await _wait_for_port(LOCALHOST, api_port, PORT_WAIT_TIMEOUT)
pid_before = proc.pid
rc, _ = await loop.run_in_executor(
None, espota2.run_ota, LOCALHOST, ota_port, None, binary_path
)
assert rc == 1, "plaintext upload without the password must fail"
await asyncio.sleep(0.5)
assert proc.returncode is None, "process died on rejected upload"
rc, _ = await loop.run_in_executor(
None, espota2.run_ota, LOCALHOST, ota_port, "hunter2", binary_path
)
assert rc == 0, "plaintext upload with the password must succeed"
await reboots.wait(1)
await _wait_for_port(LOCALHOST, api_port, PORT_WAIT_TIMEOUT)
assert proc.pid == pid_before
rc, _ = await loop.run_in_executor(
None,
functools.partial(
espota2.run_ota,
LOCALHOST,
ota_port,
None,
binary_path,
noise_psk=api_key,
),
)
assert rc == 0, "encrypted upload with the api key must succeed"
await reboots.wait(2)
await _wait_for_port(LOCALHOST, api_port, PORT_WAIT_TIMEOUT)
assert proc.pid == pid_before
@pytest.mark.asyncio
async def test_host_ota_rejects_garbage(
yaml_config: str,
+51 -4
View File
@@ -10,6 +10,7 @@ when the installed aioesphomeapi predates the noise module.
from __future__ import annotations
import base64
from collections.abc import Callable
import hashlib
import io
from pathlib import Path
@@ -109,8 +110,17 @@ 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):
if noise_negotiated and not self.offer_noise:
return # the client fails closed; nothing further arrives
if not noise_negotiated:
# A device that does not require encryption lets a plaintext
# client through
self._transfer(
lambda byte: sock.sendall(bytes([byte])),
lambda length: _recv_exact(sock, length),
lambda remaining: sock.recv(min(remaining, 4096)),
)
return
from cryptography.exceptions import InvalidTag
from noise.connection import NoiseConnection
@@ -149,6 +159,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 +183,9 @@ 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
chunk = recv_data(size - len(received))
assert chunk, "client closed mid-transfer"
received += chunk
if self.version >= espota2.OTA_VERSION_2_0:
while acked + espota2.UPLOAD_BLOCK_SIZE <= len(received) or (
len(received) == size and acked < size
@@ -240,6 +264,29 @@ def test_client_fails_closed_when_device_lacks_encryption() -> None:
device.join_and_check()
def test_plaintext_client_accepted_by_offering_device() -> None:
"""A device that offers but does not require encryption still takes a
plaintext upload from a client with no key configured."""
firmware = bytes(range(256)) * 40
device = FakeEncryptedDevice(offer_noise=True, require_noise=False)
with patch("time.sleep"):
_upload(device, firmware, None)
device.join_and_check()
assert device.received == firmware
def test_keyed_client_encrypts_with_offering_device() -> None:
"""The upload that turns on `ota: encryption:` is already encrypted when
the running firmware offers it."""
pytest.importorskip("aioesphomeapi.noise")
firmware = bytes(range(256)) * 40
device = FakeEncryptedDevice(offer_noise=True, require_noise=False)
with patch("time.sleep"):
_upload(device, firmware, PSK)
device.join_and_check()
assert device.received == firmware
def test_plaintext_client_gets_encryption_required_error() -> None:
"""A client without a key gets the device's 0x94 error message."""
device = FakeEncryptedDevice()