From 678975b6b233ef29689c78ca767c177bc53467d5 Mon Sep 17 00:00:00 2001 From: puddly <32534428+puddly@users.noreply.github.com> Date: Wed, 7 Oct 2026 14:25:28 -0400 Subject: [PATCH] [usb_host] `iManufacturer` and `iProduct` descriptor filters (#19095) Co-authored-by: J. Nick Koston Co-authored-by: Keith Burzinski --- esphome/components/usb_host/__init__.py | 59 +++++++++- esphome/components/usb_host/usb_host.h | 4 + .../components/usb_host/usb_host_client.cpp | 68 ++++++++--- esphome/components/usb_uart/usb_uart.cpp | 12 +- esphome/helpers.py | 15 +++ .../usb_host/test.esp32-s3-idf.yaml | 5 + tests/components/usb_uart/common.yaml | 19 +++ tests/unit_tests/components/test_usb_host.py | 111 ++++++++++++++++++ tests/unit_tests/test_helpers.py | 15 +++ 9 files changed, 288 insertions(+), 20 deletions(-) create mode 100644 tests/unit_tests/components/test_usb_host.py diff --git a/esphome/components/usb_host/__init__.py b/esphome/components/usb_host/__init__.py index 4abcc3a449..94532958ed 100644 --- a/esphome/components/usb_host/__init__.py +++ b/esphome/components/usb_host/__init__.py @@ -1,4 +1,7 @@ +from itertools import combinations + import esphome.codegen as cg +from esphome.components.const import CONF_MANUFACTURER from esphome.components.esp32 import ( VARIANT_ESP32H4, VARIANT_ESP32P4, @@ -15,6 +18,8 @@ from esphome.const import CONF_DEVICES, CONF_ID from esphome.core import CORE from esphome.cpp_generator import MockObj from esphome.cpp_types import Component +import esphome.final_validate as fv +from esphome.helpers import cpp_u16string_escape from esphome.types import ConfigType AUTO_LOAD = ["bytebuffer"] @@ -26,11 +31,21 @@ USBClient = usb_host_ns.class_("USBClient", Component) DOMAIN = "usb_host" CONF_VID = "vid" CONF_PID = "pid" +CONF_PRODUCT = "product" CONF_ENABLE_HUBS = "enable_hubs" CONF_MAX_TRANSFER_REQUESTS = "max_transfer_requests" CONF_MAX_PACKET_SIZE = "max_packet_size" +# VID/PID set to 0 or `None` product/manufacturer are wildcards +_FILTER_WILDCARDS = { + CONF_VID: 0, + CONF_PID: 0, + CONF_MANUFACTURER: None, + CONF_PRODUCT: None, +} + + def usb_device_schema( cls=USBClient, vid: int | None = None, pid: int | None = None ) -> cv.Schema: @@ -47,7 +62,40 @@ def usb_device_schema( schema = schema.extend({cv.Optional(CONF_PID, default=pid): cv.hex_uint16_t}) else: schema = schema.extend({cv.Required(CONF_PID): cv.hex_uint16_t}) - return schema + + return schema.extend( + { + cv.Optional(CONF_MANUFACTURER): cv.string_strict, + cv.Optional(CONF_PRODUCT): cv.string_strict, + } + ) + + +def validate_usb_clients(configs: list[ConfigType]) -> list[ConfigType]: + # Two entries overlap when no field they both constrain tells them apart + for first, second in combinations(configs, 2): + for key, wildcard in _FILTER_WILDCARDS.items(): + a = first.get(key) + b = second.get(key) + if wildcard not in (a, b) and a != b: + break + else: + raise cv.Invalid( + f"USB configs overlap: {first[CONF_ID]}, {second[CONF_ID]}" + ) + return configs + + +def _final_validate(config: ConfigType) -> ConfigType: + # Every USB client on the bus, whichever component configured it: any two could + # otherwise open the same device + clients = list(config.get(CONF_DEVICES) or ()) + clients.extend(fv.full_config.get().get("usb_uart") or ()) + validate_usb_clients(clients) + return config + + +FINAL_VALIDATE_SCHEMA = _final_validate def _set_max_packet_size(config: dict) -> dict: @@ -91,6 +139,15 @@ CONFIG_SCHEMA = cv.All( async def register_usb_client(config: ConfigType) -> MockObj: var = cg.new_Pvariable(config[CONF_ID], config[CONF_VID], config[CONF_PID]) await cg.register_component(var, config) + # UTF-16 literals, the encoding the descriptors use, so the device compares code units + if (manufacturer := config.get(CONF_MANUFACTURER)) is not None: + cg.add( + var.set_manufacturer_filter( + cg.RawExpression(cpp_u16string_escape(manufacturer)) + ) + ) + if (product := config.get(CONF_PRODUCT)) is not None: + cg.add(var.set_product_filter(cg.RawExpression(cpp_u16string_escape(product)))) return var diff --git a/esphome/components/usb_host/usb_host.h b/esphome/components/usb_host/usb_host.h index 42869fb2a6..6f92fc8538 100644 --- a/esphome/components/usb_host/usb_host.h +++ b/esphome/components/usb_host/usb_host.h @@ -143,6 +143,8 @@ class USBClient : public Component { trq_bitmask_t get_trq_in_use() const { return trq_in_use_; } bool control_transfer(uint8_t type, uint8_t request, uint16_t value, uint16_t index, const transfer_cb_t &callback, const std::vector &data = {}); + void set_manufacturer_filter(const char16_t *manufacturer) { this->manufacturer_filter_ = manufacturer; } + void set_product_filter(const char16_t *product) { this->product_filter_ = product; } // Lock-free event queue and pool for USB task to main loop communication // Must be public for access from static callbacks @@ -181,6 +183,8 @@ class USBClient : public Component { // Bit i = 1: requests_[i] is in use, Bit i = 0: requests_[i] is available // Supports multiple concurrent consumers and producers (both threads can allocate/deallocate) std::atomic trq_in_use_; + const char16_t *manufacturer_filter_{nullptr}; + const char16_t *product_filter_{nullptr}; uint16_t vid_{}; uint16_t pid_{}; }; diff --git a/esphome/components/usb_host/usb_host_client.cpp b/esphome/components/usb_host/usb_host_client.cpp index 7bc2b0a16b..3ad5fc59dc 100644 --- a/esphome/components/usb_host/usb_host_client.cpp +++ b/esphome/components/usb_host/usb_host_client.cpp @@ -11,6 +11,8 @@ #include #include #include +#include + namespace esphome::usb_host { #pragma GCC diagnostic ignored "-Wparentheses" @@ -147,21 +149,40 @@ static void usb_client_print_config_descriptor(const usb_config_desc_t *cfg_desc // Character count = (bLength - 2) / 2, max 126 chars + null terminator. static constexpr size_t DESC_STRING_BUF_SIZE = 128; -static const char *get_descriptor_string(const usb_str_desc_t *desc, std::span buffer) { - if (desc == nullptr || desc->bLength < 2) - return "(unspecified)"; - int char_count = (desc->bLength - 2) / 2; +// Folds UTF-16 to Latin-1 for logging, dropping anything that does not fit +template +static const char *utf16_to_latin1(const T *data, size_t count, std::span buffer) { char *p = buffer.data(); char *end = p + buffer.size() - 1; - for (int i = 0; i != char_count && p < end; i++) { - auto c = desc->wData[i]; - if (c < 0x100) - *p++ = static_cast(c); + for (size_t i = 0; i != count && p < end; i++) { + if (data[i] < 0x100) + *p++ = static_cast(data[i]); } *p = '\0'; return buffer.data(); } +static const char *get_descriptor_string(const usb_str_desc_t *desc, std::span buffer) { + if (desc == nullptr || desc->bLength < 2) + return "(unspecified)"; + return utf16_to_latin1(desc->wData, (desc->bLength - 2) / 2, buffer); +} + +static const char *filter_string(const char16_t *filter, std::span buffer) { + return utf16_to_latin1(filter, std::char_traits::length(filter), buffer); +} + +// Both sides are UTF-16: the descriptor by specification, the filter because code +// generation emits it as a u"" literal +static bool descriptor_string_equals(const usb_str_desc_t *desc, const char16_t *expected) { + const int char_count = (desc == nullptr || desc->bLength < 2) ? 0 : (desc->bLength - 2) / 2; + for (int i = 0; i != char_count; i++) { + if (expected[i] == u'\0' || desc->wData[i] != expected[i]) + return false; + } + return expected[char_count] == u'\0'; +} + // CALLBACK CONTEXT: USB task (called from usb_host_client_handle_events in USB task) static void client_event_cb(const usb_host_client_event_msg_t *event_msg, void *ptr) { auto *client = static_cast(ptr); @@ -302,12 +323,10 @@ void USBClient::handle_open_state_() { return; } ESP_LOGD(TAG, "Device descriptor: vid %X pid %X", desc->idVendor, desc->idProduct); - if (desc->idVendor != this->vid_ || desc->idProduct != this->pid_) { - if (this->vid_ != 0 || this->pid_ != 0) { - ESP_LOGD(TAG, "Not our device, closing"); - this->disconnect(); - return; - } + if ((this->vid_ != 0 && desc->idVendor != this->vid_) || (this->pid_ != 0 && desc->idProduct != this->pid_)) { + ESP_LOGD(TAG, "Not our device, closing"); + this->disconnect(); + return; } usb_device_info_t dev_info; err = usb_host_device_info(this->device_handle_, &dev_info); @@ -316,9 +335,21 @@ void USBClient::handle_open_state_() { this->disconnect(); return; } - this->state_ = USB_CLIENT_CONNECTED; char buf_manuf[DESC_STRING_BUF_SIZE]; char buf_product[DESC_STRING_BUF_SIZE]; + const bool manufacturer_matches = + this->manufacturer_filter_ == nullptr || + descriptor_string_equals(dev_info.str_desc_manufacturer, this->manufacturer_filter_); + const bool product_matches = + this->product_filter_ == nullptr || descriptor_string_equals(dev_info.str_desc_product, this->product_filter_); + if (!manufacturer_matches || !product_matches) { + ESP_LOGD(TAG, "Device does not match filter, closing. Manuf: %s; Prod: %s", + get_descriptor_string(dev_info.str_desc_manufacturer, buf_manuf), + get_descriptor_string(dev_info.str_desc_product, buf_product)); + this->disconnect(); + return; + } + this->state_ = USB_CLIENT_CONNECTED; char buf_serial[DESC_STRING_BUF_SIZE]; ESP_LOGD(TAG, "Device connected: Manuf: %s; Prod: %s; Serial: %s", get_descriptor_string(dev_info.str_desc_manufacturer, buf_manuf), @@ -557,6 +588,13 @@ void USBClient::dump_config() { " Vendor id %04X\n" " Product id %04X", this->vid_, this->pid_); + char buf[DESC_STRING_BUF_SIZE]; + if (this->manufacturer_filter_ != nullptr) { + ESP_LOGCONFIG(TAG, " Manufacturer %s", filter_string(this->manufacturer_filter_, buf)); + } + if (this->product_filter_ != nullptr) { + ESP_LOGCONFIG(TAG, " Product %s", filter_string(this->product_filter_, buf)); + } } // THREAD CONTEXT: Called from both USB task and main loop threads // - USB task: Immediately after transfer callback completes diff --git a/esphome/components/usb_uart/usb_uart.cpp b/esphome/components/usb_uart/usb_uart.cpp index 3113f695f6..37a5587fca 100644 --- a/esphome/components/usb_uart/usb_uart.cpp +++ b/esphome/components/usb_uart/usb_uart.cpp @@ -463,10 +463,12 @@ void USBUartTypeCdcAcm::on_connected() { void USBUartTypeCdcAcm::on_disconnected() { for (auto *channel : this->channels_) { - if (channel->cdc_dev_.in_ep != nullptr) { - usb_host_endpoint_halt(this->device_handle_, channel->cdc_dev_.in_ep->bEndpointAddress); - usb_host_endpoint_flush(this->device_handle_, channel->cdc_dev_.in_ep->bEndpointAddress); - } + // Not set up for this device: it was rejected before on_connected() ran, or it has + // fewer ports than there are channels. Nothing was claimed for it. + if (channel->cdc_dev_.in_ep == nullptr) + continue; + usb_host_endpoint_halt(this->device_handle_, channel->cdc_dev_.in_ep->bEndpointAddress); + usb_host_endpoint_flush(this->device_handle_, channel->cdc_dev_.in_ep->bEndpointAddress); if (channel->cdc_dev_.out_ep != nullptr) { usb_host_endpoint_halt(this->device_handle_, channel->cdc_dev_.out_ep->bEndpointAddress); usb_host_endpoint_flush(this->device_handle_, channel->cdc_dev_.out_ep->bEndpointAddress); @@ -494,6 +496,8 @@ void USBUartTypeCdcAcm::on_disconnected() { } } channel->initialised_.store(false); + // The descriptors these point into are freed with the device + channel->cdc_dev_ = {}; } USBClient::on_disconnected(); } diff --git a/esphome/helpers.py b/esphome/helpers.py index 3ccf0fe65a..5ec46be67e 100644 --- a/esphome/helpers.py +++ b/esphome/helpers.py @@ -201,6 +201,21 @@ def cpp_string_escape(string, encoding="utf-8"): return f'"{result}"' +def cpp_u16string_escape(string: str) -> str: + """Escape a string as a C++ u"..." literal, which the compiler encodes as UTF-16.""" + result = "" + for character in string: + code = ord(character) + if code >= 127: + # Surrogate escapes are ill-formed in C++; the compiler splits astral code points + result += f"\\U{code:08X}" + elif code < 32 or character in ("\\", '"'): + result += f"\\{code:03o}" + else: + result += character + return f'u"{result}"' + + def run_system_command(*args): import subprocess diff --git a/tests/components/usb_host/test.esp32-s3-idf.yaml b/tests/components/usb_host/test.esp32-s3-idf.yaml index 5360d1f6ff..a71140feca 100644 --- a/tests/components/usb_host/test.esp32-s3-idf.yaml +++ b/tests/components/usb_host/test.esp32-s3-idf.yaml @@ -4,3 +4,8 @@ usb_host: - id: device_1 vid: 0x1234 pid: 0x1234 + - id: device_2 + vid: 0x1234 + pid: 0xABCD + manufacturer: Example Corp + product: Example Widget diff --git a/tests/components/usb_uart/common.yaml b/tests/components/usb_uart/common.yaml index 2e41fad1a1..02500035cf 100644 --- a/tests/components/usb_uart/common.yaml +++ b/tests/components/usb_uart/common.yaml @@ -36,8 +36,10 @@ usb_uart: debug: true dummy_receiver: true debug_prefix: "[ESP_JTAG] " + # A CP2105, so it does not share 10C4:EA60 with uart_1 above - id: uart_5 type: cp210x + pid: 0xEA70 channels: - id: channel_5_1 baud_rate: 9600 @@ -66,3 +68,20 @@ usb_uart: stop_bits: 2 data_bits: 7 parity: even + # A ZBT-2 and a ZWA-2 both enumerate as 303A:4001, so only iProduct separates them + - id: uart_9 + type: cdc_acm + vid: 0x303A + pid: 0x4001 + manufacturer: Nabu Casa + product: ZBT-2 + channels: + - id: channel_9_1 + - id: uart_10 + type: cdc_acm + vid: 0x303A + pid: 0x4001 + manufacturer: Nabu Casa + product: ZWA-2 + channels: + - id: channel_10_1 diff --git a/tests/unit_tests/components/test_usb_host.py b/tests/unit_tests/components/test_usb_host.py new file mode 100644 index 0000000000..1a1df68f4d --- /dev/null +++ b/tests/unit_tests/components/test_usb_host.py @@ -0,0 +1,111 @@ +"""Tests for usb_host device matching validation.""" + +import pytest + +from esphome.components.usb_host import validate_usb_clients +import esphome.config_validation as cv +from esphome.types import ConfigType + + +@pytest.mark.parametrize( + ("first", "second"), + [ + pytest.param( + {"id": "a", "vid": 0x303A, "pid": 0x4001}, + {"id": "b", "vid": 0x303A, "pid": 0x4002}, + id="different_pid", + ), + pytest.param( + {"id": "a", "vid": 0x303A, "pid": 0x4001}, + {"id": "b", "vid": 0x1A86, "pid": 0x4001}, + id="different_vid", + ), + pytest.param( + {"id": "a", "vid": 0x303A, "pid": 0x4001, "product": "ZBT-2"}, + {"id": "b", "vid": 0x303A, "pid": 0x4001, "product": "ZWA-2"}, + id="different_product", + ), + pytest.param( + { + "id": "a", + "vid": 0x303A, + "pid": 0x4001, + "manufacturer": "Nabu Casa", + "product": "ZBT-2", + }, + { + "id": "b", + "vid": 0x303A, + "pid": 0x4001, + "manufacturer": "Espressif", + "product": "ZBT-2", + }, + id="different_manufacturer", + ), + pytest.param( + {"id": "a", "vid": 0x303A, "pid": 0, "product": "ZBT-2"}, + {"id": "b", "vid": 0x303A, "pid": 0x4001, "product": "ZWA-2"}, + id="wildcard_pid_separated_by_product", + ), + pytest.param( + {"id": "a", "vid": 0, "pid": 0x4001}, + {"id": "b", "vid": 0x303A, "pid": 0x4002}, + id="wildcard_vid_different_pid", + ), + ], +) +def test_disjoint_clients_are_accepted(first: ConfigType, second: ConfigType) -> None: + configs = [first, second] + assert validate_usb_clients(configs) is configs + + +@pytest.mark.parametrize( + ("first", "second"), + [ + pytest.param( + {"id": "a", "vid": 0x303A, "pid": 0x4001}, + {"id": "b", "vid": 0x303A, "pid": 0x4001}, + id="exact_duplicate", + ), + pytest.param( + {"id": "a", "vid": 0x303A, "pid": 0x4001}, + {"id": "b", "vid": 0x303A, "pid": 0x4001, "product": "ZBT-2"}, + id="unfiltered_shadows_filtered", + ), + pytest.param( + {"id": "a", "vid": 0x303A, "pid": 0x4001, "product": "ZBT-2"}, + {"id": "b", "vid": 0x303A, "pid": 0x4001, "product": "ZBT-2"}, + id="identical_filters", + ), + pytest.param( + {"id": "a", "vid": 0, "pid": 0}, + {"id": "b", "vid": 0x303A, "pid": 0x4001, "product": "ZBT-2"}, + id="zero_ids_match_every_device", + ), + pytest.param( + {"id": "a", "vid": 0x303A, "pid": 0}, + {"id": "b", "vid": 0x303A, "pid": 0x4001, "product": "ZBT-2"}, + id="wildcard_pid_shadows_filtered", + ), + pytest.param( + {"id": "a", "vid": 0, "pid": 0x4001}, + {"id": "b", "vid": 0x303A, "pid": 0}, + id="wildcards_on_different_fields", + ), + ], +) +def test_overlapping_clients_are_rejected( + first: ConfigType, second: ConfigType +) -> None: + with pytest.raises(cv.Invalid, match="overlap"): + validate_usb_clients([first, second]) + + +def test_every_pair_is_compared_not_just_neighbours() -> None: + configs = [ + {"id": "a", "vid": 0x303A, "pid": 0x4001, "product": "ZBT-2"}, + {"id": "b", "vid": 0x303A, "pid": 0x4002}, + {"id": "c", "vid": 0x303A, "pid": 0x4001, "product": "ZBT-2"}, + ] + with pytest.raises(cv.Invalid, match="a, c"): + validate_usb_clients(configs) diff --git a/tests/unit_tests/test_helpers.py b/tests/unit_tests/test_helpers.py index 5bdecf2fd3..8dae1e87b2 100644 --- a/tests/unit_tests/test_helpers.py +++ b/tests/unit_tests/test_helpers.py @@ -94,6 +94,21 @@ def test_cpp_string_escape(string, expected): assert actual == expected +@pytest.mark.parametrize( + "string, expected", + ( + ("foo", 'u"foo"'), + ("foo\nbar", 'u"foo\\012bar"'), + ("foo\\bar", 'u"foo\\134bar"'), + ('foo "bar"', 'u"foo \\042bar\\042"'), + ("caf\u00e9", 'u"caf\\U000000E9"'), + ("foo 🐍", 'u"foo \\U0001F40D"'), + ), +) +def test_cpp_u16string_escape(string: str, expected: str) -> None: + assert helpers.cpp_u16string_escape(string) == expected + + @pytest.mark.parametrize( "value, expected", (