diff --git a/esphome/components/nrf52/__init__.py b/esphome/components/nrf52/__init__.py index de9dc6401e..f2a7825997 100644 --- a/esphome/components/nrf52/__init__.py +++ b/esphome/components/nrf52/__init__.py @@ -732,13 +732,13 @@ def upload_program(config: ConfigType, args, host: str) -> bool: return True # Handled: PYOCD upload # Deferred imports: bleak/smpclient are heavy, only load for BLE/mcumgr paths - from .ble_logger import is_mac_address + from .ble_logger import is_ble_address from .ota import smpmgr_scan, smpmgr_upload if host == "BLE": mcumgr_device = asyncio.run(smpmgr_scan(CORE.name)) - if is_mac_address(host): + if is_ble_address(host): mcumgr_device = host if mcumgr_device: @@ -753,7 +753,7 @@ def upload_program(config: ConfigType, args, host: str) -> bool: def show_logs(config: ConfigType, args, devices: list[str]) -> bool: address = devices[0] - from .ble_logger import is_mac_address, logger_connect, logger_scan + from .ble_logger import is_ble_address, logger_connect, logger_scan if devices[0] == "BLE": ble_device = asyncio.run(logger_scan(CORE.name)) @@ -762,7 +762,7 @@ def show_logs(config: ConfigType, args, devices: list[str]) -> bool: else: return True - if is_mac_address(address): + if is_ble_address(address): asyncio.run(logger_connect(address)) return True return False diff --git a/esphome/components/nrf52/ble_logger.py b/esphome/components/nrf52/ble_logger.py index f74a49ea89..5edcddfd32 100644 --- a/esphome/components/nrf52/ble_logger.py +++ b/esphome/components/nrf52/ble_logger.py @@ -4,6 +4,7 @@ import re from typing import Final from bleak import BleakClient, BleakScanner, BLEDevice +from bleak.backends.scanner import AdvertisementData from bleak.exc import ( BleakCharacteristicNotFoundError, BleakDBusError, @@ -19,15 +20,31 @@ NUS_TX_CHAR_UUID = "6E400003-B5A3-F393-E0A9-E50E24DCCA9E" MAC_ADDRESS_PATTERN: Final = re.compile( r"([0-9A-F]{2}[:]){5}[0-9A-F]{2}$", flags=re.IGNORECASE ) +# macOS identifies BLE peripherals by a CoreBluetooth UUID instead of a MAC address +UUID_PATTERN: Final = re.compile( + r"[0-9A-F]{8}-[0-9A-F]{4}-[0-9A-F]{4}-[0-9A-F]{4}-[0-9A-F]{12}$", + flags=re.IGNORECASE, +) -def is_mac_address(value: str) -> bool: - return MAC_ADDRESS_PATTERN.match(value) +def is_ble_address(value: str) -> bool: + return bool(MAC_ADDRESS_PATTERN.match(value) or UUID_PATTERN.match(value)) + + +def _name_matches(device: BLEDevice, adv: AdvertisementData, name: str) -> bool: + # device.name can be a stale cached name on macOS; adv.local_name is what is on the air now + return name in (device.name, adv.local_name) + + +async def find_device_by_name(name: str, timeout: float = 10.0) -> BLEDevice | None: + return await BleakScanner.find_device_by_filter( + lambda device, adv: _name_matches(device, adv, name), timeout=timeout + ) async def logger_scan(name: str) -> BLEDevice | None: _LOGGER.info("Scanning bluetooth for %s...", name) - device = await BleakScanner.find_device_by_name(name) + device = await find_device_by_name(name) if not device: _LOGGER.error("%s Bluetooth LE device was not found!", name) return device diff --git a/esphome/components/nrf52/ota.py b/esphome/components/nrf52/ota.py index cafeda6478..9115e21874 100644 --- a/esphome/components/nrf52/ota.py +++ b/esphome/components/nrf52/ota.py @@ -4,7 +4,6 @@ import json import logging from pathlib import Path -from bleak import BleakScanner from bleak.exc import BleakDBusError, BleakDeviceNotFoundError from smp.exceptions import SMPBadStartDelimiter from smpclient import SMPClient @@ -23,9 +22,8 @@ from smpclient.transport.serial import SMPSerialTransport from esphome.core import EsphomeError from esphome.espota2 import ProgressBar -from .ble_logger import is_mac_address +from .ble_logger import find_device_by_name, is_ble_address -SMP_SERVICE_UUID = "8D53DC1D-1DB7-4CD3-868B-8A527460AA84" BLE_SCAN_TIMEOUT = 10.0 # seconds RESET_DELAY = 2.0 # seconds to wait before reset, allows on_end action to execute @@ -45,12 +43,12 @@ def _json_state(o: object) -> object: async def smpmgr_scan(name: str) -> str: _LOGGER.info("Scanning bluetooth for %s...", name) - for device in await BleakScanner.discover( - timeout=BLE_SCAN_TIMEOUT, service_uuids=[SMP_SERVICE_UUID] - ): - if device.name == name: - return device.address - raise EsphomeError(f"BLE device {name} with OTA service not found") + # No service filter: the SMP UUID only fits in the scan response, which macOS does not filter on. + # The SMP service is checked on connect. + device = await find_device_by_name(name, timeout=BLE_SCAN_TIMEOUT) + if device is None: + raise EsphomeError(f"BLE device {name} not found") + return device.address async def smpmgr_upload(device: str, firmware: Path) -> None: @@ -88,7 +86,7 @@ def _get_image_tlv_sha256(file: Path) -> bytes: async def _smpmgr_upload(device: str, firmware: Path) -> None: image_tlv_sha256 = _get_image_tlv_sha256(firmware) - if is_mac_address(device): + if is_ble_address(device): smp_client = SMPClient(SMPBLETransport(), device) else: smp_client = SMPClient(SMPSerialTransport(), device) @@ -105,7 +103,8 @@ async def _smpmgr_upload(device: str, firmware: Path) -> None: ) from exc raise EsphomeError(f"BLE error connecting to {device}: {exc}") from exc except SMPBLETransportException as exc: - raise EsphomeError(f"Connection error with {device}") from exc + # Reached when a device with the right name has no SMP service, among other causes + raise EsphomeError(f"Connection error with {device}: {exc}") from exc _LOGGER.info("Connected %s...", device) try: diff --git a/tests/unit_tests/test_nrf52_ota.py b/tests/unit_tests/test_nrf52_ota.py new file mode 100644 index 0000000000..c0b5e202f7 --- /dev/null +++ b/tests/unit_tests/test_nrf52_ota.py @@ -0,0 +1,121 @@ +"""Tests for nRF52 BLE device discovery and OTA transport selection.""" + +import asyncio +from pathlib import Path +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest + +from esphome.components.nrf52 import ble_logger, ota, show_logs +from esphome.core import EsphomeError + +MAC = "AA:BB:CC:DD:EE:FF" +# macOS/CoreBluetooth identifies peripherals by UUID instead of MAC address +UUID = "0FA1B2C3-D4E5-F607-1829-3A4B5C6D7E8F" + + +@pytest.mark.parametrize( + ("value", "expected"), + [ + (MAC, True), + (MAC.lower(), True), + (UUID, True), + ("/dev/ttyACM0", False), + ("COM3", False), + ], +) +def test_is_ble_address(value: str, expected: bool) -> None: + assert ble_logger.is_ble_address(value) is expected + + +def _scan_result( + monkeypatch, devices: list[tuple[str, str | None, str | None]] +) -> AsyncMock: + """Stub the bleak scan so the filter is applied to the given (address, name, local_name) devices.""" + + async def find_device_by_filter(filterfunc, timeout): + for address, name, local_name in devices: + device = SimpleNamespace(address=address, name=name) + if filterfunc(device, SimpleNamespace(local_name=local_name)): + return device + return None + + mock = AsyncMock(side_effect=find_device_by_filter) + monkeypatch.setattr(ble_logger.BleakScanner, "find_device_by_filter", mock) + return mock + + +def test_scan_matches_advertised_local_name(monkeypatch) -> None: + """On macOS device.name is cached from an earlier connection; the live name is adv.local_name.""" + _scan_result(monkeypatch, [(UUID, "stale-old-name", "my-device")]) + assert asyncio.run(ota.smpmgr_scan("my-device")) == UUID + + +def test_scan_matches_device_name(monkeypatch) -> None: + _scan_result(monkeypatch, [(MAC, "my-device", None)]) + assert asyncio.run(ota.smpmgr_scan("my-device")) == MAC + + +def test_scan_raises_when_no_device_found(monkeypatch) -> None: + _scan_result(monkeypatch, [(MAC, "someone-else", "someone-else")]) + with pytest.raises(EsphomeError, match="not found"): + asyncio.run(ota.smpmgr_scan("my-device")) + + +@pytest.mark.parametrize( + ("device", "transport"), + [("/dev/ttyACM0", "serial"), ("COM3", "serial"), (MAC, "ble"), (UUID, "ble")], +) +def test_upload_transport(monkeypatch, device: str, transport: str) -> None: + """Serial ports go to the serial transport; MAC addresses and macOS UUIDs to BLE.""" + captured: dict = {} + + class FakeClient: + def __init__(self, transport, address): + captured["transport"] = transport + captured["address"] = address + + async def connect(self) -> None: + pass + + async def disconnect(self) -> None: + pass + + monkeypatch.setattr(ota, "_get_image_tlv_sha256", lambda firmware: b"") + monkeypatch.setattr(ota, "_smpmgr_upload_connected", AsyncMock()) + monkeypatch.setattr(ota, "SMPSerialTransport", lambda: "serial") + monkeypatch.setattr(ota, "SMPBLETransport", lambda: "ble") + monkeypatch.setattr(ota, "SMPClient", FakeClient) + + asyncio.run(ota._smpmgr_upload(device, Path("firmware.bin"))) # pylint: disable=protected-access + assert captured == {"transport": transport, "address": device} + + +@pytest.mark.parametrize( + ("device", "connected_to"), + [("BLE", UUID), (MAC, MAC), (UUID, UUID), ("my-device.local", None)], +) +def test_show_logs_routing(monkeypatch, device: str, connected_to: str | None) -> None: + """A scanned device, a MAC or a UUID is handed to the BLE logger; anything else falls through.""" + scan = AsyncMock(return_value=SimpleNamespace(address=UUID)) + connect = AsyncMock(return_value=0) + monkeypatch.setattr(ble_logger, "logger_scan", scan) + monkeypatch.setattr(ble_logger, "logger_connect", connect) + + handled = show_logs(config={}, args=None, devices=[device]) + + assert handled is (connected_to is not None) + if connected_to is None: + connect.assert_not_called() + else: + connect.assert_awaited_once_with(connected_to) + + +def test_show_logs_returns_when_the_scan_finds_nothing(monkeypatch) -> None: + monkeypatch.setattr(ble_logger, "logger_scan", AsyncMock(return_value=None)) + connect = AsyncMock() + monkeypatch.setattr(ble_logger, "logger_connect", connect) + + assert show_logs(config={}, args=None, devices=["BLE"]) is True + connect.assert_not_called()