[usb_host] iManufacturer and iProduct descriptor filters (#19095)

Co-authored-by: J. Nick Koston <nick@koston.org>
Co-authored-by: Keith Burzinski <kbx81x@gmail.com>
This commit is contained in:
puddly
2026-10-07 13:25:28 -05:00
committed by GitHub
co-authored by J. Nick Koston Keith Burzinski
parent 006f5f0fc9
commit 678975b6b2
9 changed files with 288 additions and 20 deletions
+58 -1
View File
@@ -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
+4
View File
@@ -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<uint8_t> &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_bitmask_t> trq_in_use_;
const char16_t *manufacturer_filter_{nullptr};
const char16_t *product_filter_{nullptr};
uint16_t vid_{};
uint16_t pid_{};
};
+53 -15
View File
@@ -11,6 +11,8 @@
#include <cstring>
#include <atomic>
#include <span>
#include <string>
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<char, DESC_STRING_BUF_SIZE> 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<typename T>
static const char *utf16_to_latin1(const T *data, size_t count, std::span<char, DESC_STRING_BUF_SIZE> 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<char>(c);
for (size_t i = 0; i != count && p < end; i++) {
if (data[i] < 0x100)
*p++ = static_cast<char>(data[i]);
}
*p = '\0';
return buffer.data();
}
static const char *get_descriptor_string(const usb_str_desc_t *desc, std::span<char, DESC_STRING_BUF_SIZE> 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<char, DESC_STRING_BUF_SIZE> buffer) {
return utf16_to_latin1(filter, std::char_traits<char16_t>::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<USBClient *>(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
+8 -4
View File
@@ -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();
}
+15
View File
@@ -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
@@ -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
+19
View File
@@ -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
@@ -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)
+15
View File
@@ -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",
(