Compare commits

..
Author SHA1 Message Date
J. Nick Koston d4ce8cc6b4 Merge remote-tracking branch 'upstream/dev' into noise-session-resume
# Conflicts:
#	esphome/components/noise/noise.h
2026-09-08 18:19:28 +02:00
J. Nick Koston 1ce0bed3f6 [core] Share compiled binaries across modbus integration tests (#18945) 2026-09-08 18:10:45 +02:00
esphome[bot] f91486305f Bump bundled esphome-device-builder to 1.14.5 (#19040) 2026-09-08 15:21:48 +02:00
Johnandpre-commit-ci-lite[bot] <117423508+pre-commit-ci-lite[bot]@users.noreply.github.com> f191d5e0c3 [atm90e32] Verify offset calibration writes (#18701)
Co-authored-by: pre-commit-ci-lite[bot] <117423508+pre-commit-ci-lite[bot]@users.noreply.github.com>
2026-09-08 20:02:49 +12:00
Jesse Hills 227ca90aad [core] Restore the shared git hooks after post-checkout runs script/setup (#19036) 2026-09-08 20:00:08 +12:00
Jesse Hills 1700a40b7c [core] Restore the shared git hooks after post-checkout runs script/setup (#19036) 2026-09-08 19:59:55 +12:00
raykholo 28588310e7 [anova] Re-assert temperature unit on every poll cycle (#17141) 2026-09-08 18:57:33 +12:00
10a9baff74 [rf_bridge] Fix bucket sniffing with Portisch firmware (#17683)
Co-authored-by: Bryan Li <bryanli@Bryans-MacBook-Pro.local>
Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
2026-09-08 18:36:42 +12:00
Gytis a23f7bb569 [lvgl] Add missing label dependency to qrcode, keyboard and tabview (#18387) 2026-09-08 18:31:30 +12:00
J. Nick Koston 53075e4139 [core] Skip PlatformIO's private-package authorization probe (#18823) 2026-09-08 16:42:30 +12:00
Jonathan Swoboda 5722ccba37 [tuya] Build without a network component (#18948) 2026-09-08 16:36:04 +12:00
Jesse Hills 94e5c3839d [udp] Use cv.invalid for relocated packet_transport options (#19032) 2026-09-08 16:17:22 +12:00
Davide D M 574762f078 [debug] Check reboot source pref on ESP_RST_WDT and guard against empty source (#17537) 2026-09-08 13:59:05 +12:00
Pieter ViljoenandJesse Hills d34ffaf392 [ble_client] Report Established from nodes that never read services (#17920)
Co-authored-by: Jesse Hills <3060199+jesserockz@users.noreply.github.com>
2026-09-08 01:57:50 +00:00
Samuel SiebandSamuel Sieb e6aa575f2e [dallas_temp] filter 85 temp from sensor reset (#17877)
Co-authored-by: Samuel Sieb <samuel@sieb.net>
2026-09-08 13:27:49 +12:00
AndreKR 639ce609bf [logger] Fix garbled stack traces (#17939) 2026-09-08 13:13:51 +12:00
Ryan Ronnander 62eafc477d [mqtt] Restore brightness flag in light discovery (#18950) 2026-09-08 13:09:02 +12:00
mipa87andClaude 56c3361b9a [audio] Do not treat MP3_STREAM_INFO_CHANGED as a fatal decoder error (#19028)
Co-authored-by: Claude <noreply@anthropic.com>
2026-09-08 12:13:03 +12:00
mipa87Claudepre-commit-ci-lite[bot] <117423508+pre-commit-ci-lite[bot]@users.noreply.github.com>
50ca381198 [i2s_audio] Keep a start request that arrives while the speaker task stops (#19027)
Co-authored-by: Claude <noreply@anthropic.com>
Co-authored-by: pre-commit-ci-lite[bot] <117423508+pre-commit-ci-lite[bot]@users.noreply.github.com>
2026-09-08 12:10:55 +12:00
dependabot[bot] 89a56298c2 Bump prek from 0.5.1 to 0.5.2 (#19021)
Signed-off-by: dependabot[bot] <support@github.com>
2026-09-08 00:00:28 +00:00
dependabot[bot]esphome[bot] <115708604+esphome[bot]@users.noreply.github.com>Jesse Hills
390742cf9b Bump ruff from 0.16.5 to 0.16.6 (#19022)
Co-authored-by: esphome[bot] <115708604+esphome[bot]@users.noreply.github.com>
Co-authored-by: Jesse Hills <3060199+jesserockz@users.noreply.github.com>
Signed-off-by: dependabot[bot] <support@github.com>
2026-09-07 23:47:46 +00:00
Jesse Hillsandpre-commit-ci-lite[bot] <117423508+pre-commit-ci-lite[bot]@users.noreply.github.com> 5e37872da2 [ci] Sync pre-commit revs and prek version from requirements files (#19026)
Co-authored-by: pre-commit-ci-lite[bot] <117423508+pre-commit-ci-lite[bot]@users.noreply.github.com>
2026-09-08 10:32:18 +12:00
J. Nick Koston 552fcebbba Merge remote-tracking branch 'upstream/dev' into noise-session-resume 2026-09-02 11:34:05 +02:00
J. Nick Koston e90b219e00 Merge remote-tracking branch 'upstream/dev' into noise-session-resume 2026-08-26 19:49:54 -05:00
J. Nick Koston cff5cf7ad5 [noise] Trim the resume header comments 2026-08-24 19:41:53 -05:00
J. Nick Koston a9b1fc361b [noise] Document the resume constraints, bound the KDF inputs, and test PSK rotation 2026-08-24 19:39:03 -05:00
J. Nick Koston b7a5056f3e Merge remote-tracking branch 'upstream/dev' into noise-session-resume 2026-08-24 17:37:39 -05:00
J. Nick Koston a19893b168 [noise] Inline the resume MAC and key helpers so try_accept calls the KDF directly 2026-08-24 16:55:30 -05:00
J. Nick Koston 2344929dae [noise] Keep the resume KDF labels in PROGMEM on ESP8266 2026-08-24 16:50:28 -05:00
J. Nick Koston 46de7335b2 [noise] Simplify the resume cache accept path and guard the sensitive message set 2026-08-24 16:44:47 -05:00
J. Nick Koston 1cb2136c38 [noise] Trim resume flash: one KDF, no discard state, packed ticket 2026-08-24 16:31:46 -05:00
J. Nick Koston d1f7e71efb [noise] Trim resume flash usage and never dump the ticket secret 2026-08-24 15:53:33 -05:00
J. Nick Koston cf97e4c4e8 [noise] Add a session resume integration test 2026-08-24 15:41:23 -05:00
J. Nick Koston 4e2ff94588 [noise] Keep two resume tickets 2026-08-24 14:39:15 -05:00
J. Nick Koston c2118778ee [noise] Move NoiseResumeTicket to message id 152 2026-08-24 14:31:36 -05:00
J. Nick Koston c655a2a442 [noise] Use NOLINT for the label memcpy 2026-08-24 14:31:36 -05:00
J. Nick Koston b8a79a26fb [noise] Trim comments 2026-08-24 14:31:35 -05:00
J. Nick Koston 34caf86839 [noise] Fix clang-tidy findings 2026-08-24 14:31:35 -05:00
J. Nick Koston 1eb4cff6b6 [noise] Simplify the resume implementation 2026-08-24 14:31:35 -05:00
J. Nick Koston b60950da76 [noise] Add known answer and cache tests for session resume 2026-08-24 14:31:35 -05:00
J. Nick Koston 7e6db4d335 [noise] Add session resume to the api noise transport 2026-08-24 14:31:35 -05:00
102 changed files with 6414 additions and 5785 deletions
+11 -2
View File
@@ -244,11 +244,20 @@ jobs:
steps:
- name: Check out code from GitHub
uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
- name: Read prek version from requirements_test.txt
id: prek
# requirements_test.txt is the only place the version is pinned, so a
# Dependabot bump there is picked up here without a second edit.
run: |
if ! version=$(sed -nE 's/^prek==([^[:space:]#]+).*/\1/p' requirements_test.txt) || [ -z "$version" ]; then
echo "::error::No prek== pin found in requirements_test.txt."
exit 1
fi
echo "version=$version" >> "$GITHUB_OUTPUT"
- name: Run prek
uses: j178/prek-action@4e14d07f9231acabce116ccfca13b13dd9755ece # v3.0.0
with:
# Keep in sync with requirements_test.txt.
prek-version: "0.4.11"
prek-version: ${{ steps.prek.outputs.version }}
# This job only runs on pull requests, so nothing ever populates
# the cache on dev. Every run would miss and then write a per-pull
# request copy, which is what the old seed-cache job existed to
@@ -0,0 +1,94 @@
# Keeps pre-commit hook revs in sync with the requirements files.
#
# Dependabot only bumps the pins in requirements*.txt. Some of those tools
# are pinned again as hook revs in .pre-commit-config.yaml. This workflow
# runs script/sync_dependency_versions.py against the pull request branch
# and pushes a commit with the revs updated.
name: Sync dependency versions
on:
# pull_request_target rather than pull_request so the App secret is
# available on Dependabot pull requests (pull_request runs opened by
# Dependabot only see Dependabot secrets). The job below only touches
# branches in this repository and only ever executes the script from the
# base branch checkout, so fork code never runs with the token.
pull_request_target:
types: [opened, synchronize, reopened]
paths:
- requirements_dev.txt
- requirements_test.txt
- .pre-commit-config.yaml
- script/sync_dependency_versions.py
# The push to the pull request branch uses the App token minted below, so
# the workflow's GITHUB_TOKEN does not need any scopes.
permissions: {}
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number }}
cancel-in-progress: true
jobs:
sync:
name: Sync pinned versions
runs-on: ubuntu-latest
# Same-repository branches only: a push to a fork is not possible with
# this token, and it keeps untrusted heads out of a privileged job.
if: >-
github.repository == 'esphome/esphome'
&& github.event.pull_request.head.repo.full_name == github.repository
steps:
- name: Generate a token
id: generate-token
uses: actions/create-github-app-token@bcd2ba49218906704ab6c1aa796996da409d3eb1 # v3.2.0
with:
client-id: ${{ vars.ESPHOME_GITHUB_APP_CLIENT_ID }}
private-key: ${{ secrets.ESPHOME_GITHUB_APP_PRIVATE_KEY }}
# A push made with the workflow's own GITHUB_TOKEN would not start
# CI on the new commit; a push with the App token does.
permission-contents: write # git push of the sync commit to the pull request branch
- name: Check out base branch
# Provides the script that runs below. Deliberately the base branch
# so the pull request cannot change what executes here.
uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
with:
ref: ${{ github.event.pull_request.base.sha }}
persist-credentials: false
- name: Check out pull request branch
# No allow-unsafe-pr-checkout here on purpose: checkout v7 only
# refuses heads that live in a different repository, and the job
# condition above already limits runs to same-repository branches.
# Leaving it off keeps that refusal as a backstop for fork heads.
uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
with:
ref: ${{ github.event.pull_request.head.ref }}
path: pull-request
token: ${{ steps.generate-token.outputs.token }}
- name: Set up Python
uses: actions/setup-python@5fda3b95a4ea91299a34e894583c3862153e4b97 # v7.0.0
with:
python-version: "3.12"
- name: Install yamlrocks
# The script edits YAML through yamlrocks. Take the pin from the
# base branch requirements so this workflow has no copy of its own.
run: pip install "$(grep -E '^yamlrocks==' requirements_test.txt | cut -d'#' -f1)"
- name: Sync pinned versions
run: python script/sync_dependency_versions.py --root pull-request
- name: Push changes
working-directory: pull-request
run: |
if git diff --quiet; then
echo "All pinned versions already match the requirements files."
exit 0
fi
git config user.name "esphome[bot]"
git config user.email "115708604+esphome[bot]@users.noreply.github.com"
git commit -am "Sync pinned tool versions with requirements files"
git push
+2 -3
View File
@@ -1,7 +1,6 @@
---
# See https://pre-commit.com for more information
# See https://pre-commit.com/hooks.html for more hooks
ci:
autoupdate_commit_msg: 'pre-commit: autoupdate'
autoupdate_schedule: off # Disabled until ruff versions are synced between deps and pre-commit
@@ -11,7 +10,7 @@ ci:
repos:
- repo: https://github.com/astral-sh/ruff-pre-commit
# Ruff version.
rev: v0.16.3
rev: v0.16.6
hooks:
# Run the linter.
- id: ruff
@@ -42,7 +41,7 @@ repos:
- id: pyupgrade
args: [--py312-plus]
- repo: https://github.com/adrienverge/yamllint.git
rev: v1.37.1
rev: v1.38.0
hooks:
- id: yamllint
exclude: ^(\.clang-format|\.clang-tidy)$
+1 -1
View File
@@ -840,7 +840,7 @@ file does, and it is the authority when they disagree. The most useful starting
cv.rename_key(
CONF_OLD_KEY, CONF_NEW_KEY, removed_in="2026.6.0", component="my_component"
),
cv.Schema({ ... }),
cv.Schema({...}),
)
```
For other deprecations, warn manually during validation:
+1 -1
View File
@@ -22,7 +22,7 @@ RUN \
-r /requirements.txt
# Install the ESPHome Device Builder dashboard.
RUN uv pip install --no-cache-dir esphome-device-builder==1.14.4
RUN uv pip install --no-cache-dir esphome-device-builder==1.14.5
RUN \
platformio settings set enable_telemetry No \
+1 -3
View File
@@ -23,9 +23,7 @@ from esphome.util import safe_print
if TYPE_CHECKING:
from collections.abc import Callable
from aioesphomeapi.api_pb2 import (
SubscribeLogsResponse, # pylint: disable=no-name-in-module
)
from aioesphomeapi.api_pb2 import SubscribeLogsResponse # pylint: disable=no-name-in-module
_LOGGER = logging.getLogger(__name__)
+51 -56
View File
@@ -13,7 +13,7 @@ void Anova::dump_config() { LOG_CLIMATE("", "Anova BLE Cooker", this); }
void Anova::setup() {
this->codec_ = make_unique<AnovaCodec>();
this->current_request_ = 0;
this->poll_step_ = PollStep::IDLE;
}
void Anova::loop() {
@@ -22,6 +22,15 @@ void Anova::loop() {
this->disable_loop();
}
void Anova::write_request_(AnovaPacket *pkt) {
auto status =
esp_ble_gattc_write_char(this->parent_->get_gattc_if(), this->parent_->get_conn_id(), this->char_handle_,
pkt->length, pkt->data, ESP_GATT_WRITE_TYPE_NO_RSP, ESP_GATT_AUTH_REQ_NONE);
if (status) {
ESP_LOGW(TAG, "[%s] esp_ble_gattc_write_char failed, status=%d", this->parent_->address_str(), status);
}
}
void Anova::control(const ClimateCall &call) {
auto mode_val = call.get_mode();
if (mode_val.has_value()) {
@@ -38,22 +47,11 @@ void Anova::control(const ClimateCall &call) {
ESP_LOGW(TAG, "Unsupported mode: %d", mode);
return;
}
auto status =
esp_ble_gattc_write_char(this->parent_->get_gattc_if(), this->parent_->get_conn_id(), this->char_handle_,
pkt->length, pkt->data, ESP_GATT_WRITE_TYPE_NO_RSP, ESP_GATT_AUTH_REQ_NONE);
if (status) {
ESP_LOGW(TAG, "[%s] esp_ble_gattc_write_char failed, status=%d", this->parent_->address_str(), status);
}
this->write_request_(pkt);
}
auto target_temp = call.get_target_temperature();
if (target_temp.has_value()) {
auto *pkt = this->codec_->get_set_target_temp_request(*target_temp);
auto status =
esp_ble_gattc_write_char(this->parent_->get_gattc_if(), this->parent_->get_conn_id(), this->char_handle_,
pkt->length, pkt->data, ESP_GATT_WRITE_TYPE_NO_RSP, ESP_GATT_AUTH_REQ_NONE);
if (status) {
ESP_LOGW(TAG, "[%s] esp_ble_gattc_write_char failed, status=%d", this->parent_->address_str(), status);
}
this->write_request_(this->codec_->get_set_target_temp_request(*target_temp));
}
}
@@ -62,6 +60,7 @@ void Anova::gattc_event_handler(esp_gattc_cb_event_t event, esp_gatt_if_t gattc_
case ESP_GATTC_DISCONNECT_EVT: {
this->current_temperature = NAN;
this->target_temperature = NAN;
this->poll_step_ = PollStep::IDLE;
this->publish_state();
break;
}
@@ -83,8 +82,8 @@ void Anova::gattc_event_handler(esp_gattc_cb_event_t event, esp_gatt_if_t gattc_
}
case ESP_GATTC_REG_FOR_NOTIFY_EVT: {
this->node_state = espbt::ClientState::ESTABLISHED;
this->current_request_ = 0;
this->update();
this->poll_step_ = PollStep::IDLE;
this->update(); // begin the first poll cycle immediately
break;
}
case ESP_GATTC_NOTIFY_EVT: {
@@ -101,33 +100,30 @@ void Anova::gattc_event_handler(esp_gattc_cb_event_t event, esp_gatt_if_t gattc_
this->mode = this->codec_->running_ ? climate::CLIMATE_MODE_HEAT : climate::CLIMATE_MODE_OFF;
}
if (this->codec_->has_unit()) {
this->fahrenheit_ = (this->codec_->unit_ == 'f');
ESP_LOGD(TAG, "Anova units is %s", this->fahrenheit_ ? "fahrenheit" : "celsius");
this->current_request_++;
ESP_LOGD(TAG, "Anova units is %s", (this->codec_->unit_ == 'f') ? "fahrenheit" : "celsius");
}
this->publish_state();
if (this->current_request_ > 1) {
AnovaPacket *pkt = nullptr;
switch (this->current_request_++) {
case 2:
pkt = this->codec_->get_read_target_temp_request();
break;
case 3:
pkt = this->codec_->get_read_current_temp_request();
break;
default:
this->current_request_ = 1;
break;
}
if (pkt != nullptr) {
auto status =
esp_ble_gattc_write_char(this->parent_->get_gattc_if(), this->parent_->get_conn_id(), this->char_handle_,
pkt->length, pkt->data, ESP_GATT_WRITE_TYPE_NO_RSP, ESP_GATT_AUTH_REQ_NONE);
if (status) {
ESP_LOGW(TAG, "[%s] esp_ble_gattc_write_char failed, status=%d", this->parent_->address_str(), status);
}
}
// Advance the poll cycle to its next request based on the reply we got.
switch (this->poll_step_) {
case PollStep::SET_UNIT:
this->poll_step_ = PollStep::STATUS;
this->write_request_(this->codec_->get_read_device_status_request());
break;
case PollStep::STATUS:
this->poll_step_ = PollStep::TARGET;
this->write_request_(this->codec_->get_read_target_temp_request());
break;
case PollStep::TARGET:
this->poll_step_ = PollStep::CURRENT;
this->write_request_(this->codec_->get_read_current_temp_request());
break;
case PollStep::CURRENT:
this->poll_step_ = PollStep::IDLE; // full cycle complete
break;
default:
// A reply to an ad-hoc control() write, outside a managed cycle.
break;
}
break;
}
@@ -136,27 +132,26 @@ void Anova::gattc_event_handler(esp_gattc_cb_event_t event, esp_gatt_if_t gattc_
}
}
void Anova::set_unit_of_measurement(const char *unit) { this->fahrenheit_ = !strncmp(unit, "f", 1); }
void Anova::set_unit_of_measurement(const char *unit) { this->want_fahrenheit_ = !strncmp(unit, "f", 1); }
void Anova::update() {
if (this->node_state != espbt::ClientState::ESTABLISHED)
return;
if (this->current_request_ < 2) {
AnovaPacket *pkt;
if (this->current_request_ == 0) {
pkt = this->codec_->get_set_unit_request(this->fahrenheit_ ? 'f' : 'c');
} else {
pkt = this->codec_->get_read_device_status_request();
}
auto status =
esp_ble_gattc_write_char(this->parent_->get_gattc_if(), this->parent_->get_conn_id(), this->char_handle_,
pkt->length, pkt->data, ESP_GATT_WRITE_TYPE_NO_RSP, ESP_GATT_AUTH_REQ_NONE);
if (status) {
ESP_LOGW(TAG, "[%s] esp_ble_gattc_write_char failed, status=%d", this->parent_->address_str(), status);
}
this->current_request_++;
if (this->poll_step_ != PollStep::IDLE) {
// The previous cycle never finished within a full polling interval -- a
// reply was missed or a write failed. Restart the cycle rather than stall;
// the polling interval itself acts as the timeout. A late reply from the
// abandoned cycle is harmless: state decoding happens on every notify
// regardless of step, and each notify sends at most one follow-up request.
ESP_LOGW(TAG, "[%s] Poll cycle incomplete (step %u); restarting cycle", this->parent_->address_str(),
static_cast<uint8_t>(this->poll_step_));
}
// Re-assert the configured unit at the start of every poll cycle, then fall
// through the status/temperature reads via the notification handler. Always
// command the configured unit (want_fahrenheit_) -- never the last value the
// device reported, or a drift to 'c' would lock itself in.
this->poll_step_ = PollStep::SET_UNIT;
this->write_request_(this->codec_->get_set_unit_request(this->want_fahrenheit_ ? 'f' : 'c'));
}
} // namespace esphome::anova
+11 -2
View File
@@ -37,11 +37,20 @@ class Anova final : public climate::Climate, public esphome::ble_client::BLEClie
void set_unit_of_measurement(const char *unit);
protected:
// A poll cycle re-asserts the configured unit, then reads device state.
// Re-asserting every cycle prevents the cooker from silently reverting to
// its default (Celsius); previously the unit was only set once on
// connection, so a drift persisted (and corrupted the F/C interpretation of
// subsequent readings) until the BLE link was re-established.
enum class PollStep : uint8_t { SET_UNIT, STATUS, TARGET, CURRENT, IDLE };
void write_request_(AnovaPacket *pkt);
std::unique_ptr<AnovaCodec> codec_;
void control(const climate::ClimateCall &call) override;
uint16_t char_handle_;
uint8_t current_request_;
bool fahrenheit_;
bool want_fahrenheit_{true}; // configured target unit; never overwritten by device replies
PollStep poll_step_{PollStep::IDLE};
};
} // namespace esphome::anova
+14
View File
@@ -893,6 +893,20 @@ message NoiseEncryptionSetKeyResponse {
bool success = 1;
}
// Single-use session resume ticket, sent unsolicited by the device after a
// Noise connection authenticates. A client presents it in the ClientHello of
// its next connection to skip the curve25519 handshake; the device then
// issues a fresh ticket on that connection. Never sent on plaintext
// connections. Clients that do not understand it drop it silently.
// Contents are secret; the device generator redacts this message from dump_to
message NoiseResumeTicket {
option (id) = 152;
option (source) = SOURCE_SERVER;
option (ifdef) = "USE_API_NOISE";
bytes ticket = 1; // session_id(8) || secret(32)
}
// ==================== HOMEASSISTANT.SERVICE ====================
message SubscribeHomeassistantServicesRequest {
option (id) = 34;
+25 -6
View File
@@ -1779,6 +1779,9 @@ void APIConnection::complete_authentication_() {
this->send_time_request();
}
#endif
#ifdef USE_API_NOISE
this->send_resume_ticket_();
#endif
#ifdef USE_ZWAVE_PROXY
if (zwave_proxy::global_zwave_proxy != nullptr) {
zwave_proxy::global_zwave_proxy->api_connection_authenticated(this);
@@ -1786,6 +1789,27 @@ void APIConnection::complete_authentication_() {
#endif
}
#ifdef USE_API_NOISE
void APIConnection::send_resume_ticket_() {
#ifdef USE_API_PLAINTEXT
// Only encrypted transports get a ticket: on dual-mode builds a plaintext
// connection has no frame footer
if (this->helper_->frame_footer_size() == 0) {
return;
}
#endif
noise::ResumeTicket ticket;
if (!this->parent_->get_noise_ctx().resume_cache().issue(ticket)) {
return;
}
NoiseResumeTicket msg;
msg.set_ticket(reinterpret_cast<const uint8_t *>(&ticket), sizeof(ticket));
// A dropped ticket is harmless: the client does a full handshake next time
static_cast<void>(this->send_message(msg));
noise_clean(&ticket, sizeof(ticket));
}
#endif
bool APIConnection::send_hello_response_(const HelloRequest &msg) {
// Copy client name with truncation if needed (set_client_name handles truncation)
this->helper_->set_client_name(msg.client_info.c_str(), msg.client_info.size());
@@ -2255,12 +2279,7 @@ bool APIConnection::send_message_(uint32_t payload_size, uint16_t message_type,
// Capacity reserved above, cannot fail
(void) shared_buf.resize(write_start + payload_size);
ProtoWriteBuffer buffer{&shared_buf, write_start};
uint8_t *end = encode_fn(msg, buffer PROTO_ENCODE_DEBUG_INIT(&shared_buf));
#ifdef ESPHOME_DEBUG_API
assert(end == shared_buf.data() + shared_buf.size());
#else
(void) end;
#endif
encode_fn(msg, buffer PROTO_ENCODE_DEBUG_INIT(&shared_buf));
return this->send_buffer(ProtoWriteBuffer{&shared_buf}, message_type);
}
// encode_to_buffer is defined inline in api_connection.h (ESPHOME_ALWAYS_INLINE)
+29 -6
View File
@@ -345,7 +345,11 @@ class APIConnection final : public APIServerConnectionBase {
/// Returns false as soon as the TCP buffer is full. Marked nodiscard so we
/// have no silent failures: every caller must handle (or log) a refusal.
template<typename T> [[nodiscard]] bool send_message(const T &msg) {
return this->send_message_(T::calc_size_msg(&msg), T::MESSAGE_TYPE, &T::encode_msg, &msg);
if constexpr (T::ESTIMATED_SIZE == 0) {
return this->send_message_(0, T::MESSAGE_TYPE, &encode_msg_noop, &msg);
} else {
return this->send_message_(msg.calculate_size(), T::MESSAGE_TYPE, &proto_encode_msg<T>, &msg);
}
}
/// Clear the shared write buffer and reserve space for the first message.
@@ -377,6 +381,11 @@ class APIConnection final : public APIServerConnectionBase {
// Helper function to handle authentication completion
void complete_authentication_();
#ifdef USE_API_NOISE
// Issue a fresh single-use session resume ticket over the encrypted channel
void send_resume_ticket_();
#endif
// Pattern B helpers: send response and return success/failure
bool send_hello_response_(const HelloRequest &msg);
bool send_disconnect_response_();
@@ -401,6 +410,16 @@ class APIConnection final : public APIServerConnectionBase {
void process_state_subscriptions_();
#endif
// Size thunk — converts void* back to concrete type for direct calculate_size() call
template<typename T> static uint32_t calc_size(const void *msg) {
return static_cast<const T *>(msg)->calculate_size();
}
// Shared no-op encode thunk for empty messages (ESTIMATED_SIZE == 0)
static uint8_t *encode_msg_noop(const void *, ProtoWriteBuffer &buf PROTO_ENCODE_DEBUG_PARAM) {
return buf.get_pos();
}
// Non-template buffer management for send_message
bool send_message_(uint32_t payload_size, uint16_t message_type, MessageEncodeFn encode_fn, const void *msg);
@@ -419,7 +438,11 @@ class APIConnection final : public APIServerConnectionBase {
// Hot paths (state/info) go through fill_and_encode_entity_state/info instead.
// batch_message_type_ is already set by dispatch_message_ before reaching here.
template<typename T> static uint16_t encode_message_to_buffer(T &msg, APIConnection *conn, uint32_t remaining_size) {
return encode_to_buffer_slow(T::calc_size_msg(&msg), &T::encode_msg, &msg, conn, remaining_size);
if constexpr (T::ESTIMATED_SIZE == 0) {
return encode_to_buffer_slow(0, &encode_msg_noop, &msg, conn, remaining_size);
} else {
return encode_to_buffer_slow(msg.calculate_size(), &proto_encode_msg<T>, &msg, conn, remaining_size);
}
}
// Non-template core — fills state fields and encodes
@@ -431,7 +454,7 @@ class APIConnection final : public APIServerConnectionBase {
template<typename T>
static uint16_t fill_and_encode_entity_state(EntityBase *entity, T &msg, APIConnection *conn,
uint32_t remaining_size) {
return fill_and_encode_entity_state(entity, msg, &T::calc_size_msg, &T::encode_msg, conn, remaining_size);
return fill_and_encode_entity_state(entity, msg, &calc_size<T>, &proto_encode_msg<T>, conn, remaining_size);
}
// Non-template core — fills info fields, allocates buffers, and encodes
@@ -443,7 +466,7 @@ class APIConnection final : public APIServerConnectionBase {
template<typename T>
static uint16_t fill_and_encode_entity_info(EntityBase *entity, T &msg, APIConnection *conn,
uint32_t remaining_size) {
return fill_and_encode_entity_info(entity, msg, &T::calc_size_msg, &T::encode_msg, conn, remaining_size);
return fill_and_encode_entity_info(entity, msg, &calc_size<T>, &proto_encode_msg<T>, conn, remaining_size);
}
// Non-template core — fills device_class, then delegates to fill_and_encode_entity_info
@@ -457,8 +480,8 @@ class APIConnection final : public APIServerConnectionBase {
static uint16_t fill_and_encode_entity_info_with_device_class(EntityBase *entity, T &msg,
StringRef &device_class_field, APIConnection *conn,
uint32_t remaining_size) {
return fill_and_encode_entity_info_with_device_class(entity, msg, device_class_field, &T::calc_size_msg,
&T::encode_msg, conn, remaining_size);
return fill_and_encode_entity_info_with_device_class(entity, msg, device_class_field, &calc_size<T>,
&proto_encode_msg<T>, conn, remaining_size);
}
#ifdef USE_VOICE_ASSISTANT
@@ -46,13 +46,7 @@ inline uint16_t ESPHOME_ALWAYS_INLINE APIConnection::encode_to_buffer(uint32_t c
return 0;
}
ProtoWriteBuffer buffer{&shared_buf, shared_buf.size() - calculated_size};
uint8_t *end = encode_fn(msg, buffer PROTO_ENCODE_DEBUG_INIT(&shared_buf));
#ifdef ESPHOME_DEBUG_API
// A body that writes fewer bytes than calculate_size() promised would ship stale buffer bytes
assert(end == shared_buf.data() + shared_buf.size());
#else
(void) end;
#endif
encode_fn(msg, buffer PROTO_ENCODE_DEBUG_INIT(&shared_buf));
return total_calculated_size;
}
@@ -271,8 +271,8 @@ APIError APINoiseFrameHelper::state_action_client_hello_() {
if (aerr != APIError::OK) {
return handle_handshake_frame_error_(aerr);
}
// ignore contents, may be used in future for flags
// Resize for: existing prologue + 2 size bytes + frame data
// Contents are extension flags (today: the resume offer); mixed into the
// prologue either way. Resize for: existing prologue + 2 size bytes + frame data
size_t old_size = this->prologue_.size();
size_t rx_size = this->rx_buf_.size();
if (!this->prologue_.resize(old_size + 2 + rx_size)) [[unlikely]] {
@@ -289,6 +289,8 @@ APIError APINoiseFrameHelper::state_action_client_hello_() {
return APIError::OK;
}
APIError APINoiseFrameHelper::state_action_server_hello_() {
// A verified resume offer (still in rx_buf_ from the client hello step)
// replaces the whole handshake; any failure falls back to the full one.
// send server hello
const auto &name = App.get_name();
char mac[MAC_ADDRESS_BUFFER_SIZE];
@@ -302,7 +304,9 @@ APIError APINoiseFrameHelper::state_action_server_hello_() {
// 1 (proto) + name (max ESPHOME_DEVICE_NAME_MAX_LEN) + 1 (name null)
// + mac (MAC_ADDRESS_BUFFER_SIZE - 1) + 1 (mac null)
constexpr size_t max_msg_size = 1 + ESPHOME_DEVICE_NAME_MAX_LEN + 1 + MAC_ADDRESS_BUFFER_SIZE;
// + optional resume accept extension
constexpr size_t max_msg_size =
1 + ESPHOME_DEVICE_NAME_MAX_LEN + 1 + MAC_ADDRESS_BUFFER_SIZE + noise::RESUME_ACCEPT_SIZE;
uint8_t msg[max_msg_size];
// chosen proto
@@ -313,16 +317,32 @@ APIError APINoiseFrameHelper::state_action_server_hello_() {
// node mac, terminated by null byte
std::memcpy(msg + mac_offset, mac, MAC_ADDRESS_BUFFER_SIZE);
// The accept extension, if any, is written straight after the mac
size_t ext_len = this->ctx_.resume_cache().try_accept(
this->rx_buf_.data(), this->rx_buf_.size(), this->prologue_.data(), this->prologue_.size(), msg + total_size,
sizeof(msg) - total_size, send_cipher_, recv_cipher_);
bool resume = ext_len != 0;
total_size += ext_len;
APIError aerr = write_frame_(msg, total_size);
if (aerr != APIError::OK)
return aerr;
// start handshake
aerr = init_handshake_();
if (aerr != APIError::OK)
return aerr;
state_ = State::HANDSHAKE;
if (resume) {
// A resuming client waits for this hello instead of pipelining
// handshake message 1, so the transport is ready now
this->frame_footer_size_ = noise_cipherstate_get_mac_length(this->send_cipher_);
HELPER_LOG("Session resumed!");
state_ = State::DATA;
} else {
aerr = init_handshake_();
if (aerr != APIError::OK)
return aerr;
state_ = State::HANDSHAKE;
}
// init_handshake_ copied the prologue into the handshake state; the resume
// path is done with it too
this->prologue_.release();
return APIError::OK;
}
APIError APINoiseFrameHelper::state_action_handshake_() {
@@ -552,8 +572,6 @@ APIError APINoiseFrameHelper::init_handshake_() {
APIError aerr = handle_noise_error_(err, LOG_STR("noise_handshake_init"), APIError::HANDSHAKESTATE_SETUP_FAILED);
if (aerr != APIError::OK)
return aerr;
// init copies the prologue into the handshakestate, so we can get rid of it now
prologue_.release();
return APIError::OK;
}
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+4
View File
@@ -1393,6 +1393,10 @@ const char *NoiseEncryptionSetKeyResponse::dump_to(DumpBuffer &out) const {
dump_field(out, ESPHOME_PSTR("success"), this->success);
return out.c_str();
}
const char *NoiseResumeTicket::dump_to(DumpBuffer &out) const {
out.append_p(ESPHOME_PSTR("NoiseResumeTicket {}"));
return out.c_str();
}
#endif
#ifdef USE_API_HOMEASSISTANT_SERVICES
const char *HomeassistantServiceMap::dump_to(DumpBuffer &out) const {
+2 -3
View File
@@ -433,9 +433,8 @@ void APIServer::send_homeassistant_action(const HomeassistantActionRequest &call
// Home Assistant subscribes to actions shortly *after* authenticating, so actions
// fired right at connection time (on_client_connected, on_time_sync, ...) can
// arrive before the subscription and are lost - warn instead of failing silently.
ESP_LOGW(TAG, "Home Assistant %s '%.*s' dropped; %s",
call.is_event ? LOG_STR_LITERAL("event") : LOG_STR_LITERAL("action"),
static_cast<int>(call.service.size()), call.service.empty() ? "" : call.service.c_str(),
ESP_LOGW(TAG, "Home Assistant %s '%s' dropped; %s",
call.is_event ? LOG_STR_LITERAL("event") : LOG_STR_LITERAL("action"), call.service.c_str(),
this->is_connected() ? LOG_STR_LITERAL("client has not subscribed to actions (yet)")
: LOG_STR_LITERAL("no client connected"));
}
+57 -58
View File
@@ -214,74 +214,73 @@ void ProtoDecodableMessage::decode(const uint8_t *buffer, size_t length) {
const uint8_t *ptr = buffer;
const uint8_t *end = buffer + length;
// Single-byte varints dominate, so that case advances the cursor inline.
auto read_varint = [&](proto_varint_value_t &value) ESPHOME_ALWAYS_INLINE {
if (ptr == end)
return false;
if (*ptr < 0x80) [[likely]] {
value = *ptr++;
return true;
}
auto res = ProtoVarInt::parse_non_empty(ptr, end - ptr);
if (!res.has_value())
return false;
value = res.value;
ptr += res.consumed;
return true;
};
while (ptr < end) {
proto_varint_value_t tag_value;
if (!read_varint(tag_value)) {
// Parse field header - ptr < end guarantees len >= 1
auto res = ProtoVarInt::parse_non_empty(ptr, end - ptr);
if (!res.has_value()) {
ESP_LOGV(TAG, "Invalid field start at offset %ld", (long) (ptr - buffer));
return;
}
uint32_t tag = static_cast<uint32_t>(tag_value);
uint32_t tag = static_cast<uint32_t>(res.value);
uint32_t field_type = tag & WIRE_TYPE_MASK;
// Length-delimited fields move this past the length prefix
const uint8_t *data = ptr;
proto_varint_value_t scalar;
uint32_t field_id = tag >> 3;
ptr += res.consumed;
if (field_type == WIRE_TYPE_VARINT) [[likely]] {
if (!read_varint(scalar)) {
ESP_LOGV(TAG, "Invalid VarInt at offset %ld", (long) (ptr - buffer));
return;
}
} else {
switch (field_type) {
case WIRE_TYPE_LENGTH_DELIMITED: {
proto_varint_value_t length_value;
if (!read_varint(length_value)) {
ESP_LOGV(TAG, "Invalid Length Delimited at offset %ld", (long) (ptr - buffer));
return;
}
uint32_t field_length = static_cast<uint32_t>(length_value);
if (field_length > static_cast<size_t>(end - ptr)) {
ESP_LOGV(TAG, "Out-of-bounds Length Delimited at offset %ld", (long) (ptr - buffer));
return;
}
data = ptr;
scalar = field_length;
ptr += field_length;
break;
}
case WIRE_TYPE_FIXED32: {
if (end - ptr < 4) {
ESP_LOGV(TAG, "Out-of-bounds Fixed32-bit at offset %ld", (long) (ptr - buffer));
return;
}
// Byte loads instead of memcpy: ESP-IDF passes -fno-builtin-memcpy, which made this a call
scalar = encode_uint32(ptr[3], ptr[2], ptr[1], ptr[0]);
ptr += 4;
break;
}
default:
ESP_LOGV(TAG, "Invalid field type %" PRIu32 " at offset %ld", field_type, (long) (ptr - buffer));
switch (field_type) {
case WIRE_TYPE_VARINT: { // VarInt
res = ProtoVarInt::parse(ptr, end - ptr);
if (!res.has_value()) {
ESP_LOGV(TAG, "Invalid VarInt at offset %ld", (long) (ptr - buffer));
return;
}
if (!this->decode_varint(field_id, res.value)) {
ESP_LOGV(TAG, "Cannot decode VarInt field %" PRIu32 " with value %" PRIu64 "!", field_id,
static_cast<uint64_t>(res.value));
}
ptr += res.consumed;
break;
}
case WIRE_TYPE_LENGTH_DELIMITED: { // Length-delimited
res = ProtoVarInt::parse(ptr, end - ptr);
if (!res.has_value()) {
ESP_LOGV(TAG, "Invalid Length Delimited at offset %ld", (long) (ptr - buffer));
return;
}
uint32_t field_length = static_cast<uint32_t>(res.value);
ptr += res.consumed;
if (field_length > static_cast<size_t>(end - ptr)) {
ESP_LOGV(TAG, "Out-of-bounds Length Delimited at offset %ld", (long) (ptr - buffer));
return;
}
if (!this->decode_length(field_id, ProtoLengthDelimited(ptr, field_length))) {
ESP_LOGV(TAG, "Cannot decode Length Delimited field %" PRIu32 "!", field_id);
}
ptr += field_length;
break;
}
case WIRE_TYPE_FIXED32: { // 32-bit
if (end - ptr < 4) {
ESP_LOGV(TAG, "Out-of-bounds Fixed32-bit at offset %ld", (long) (ptr - buffer));
return;
}
uint32_t val;
#if __BYTE_ORDER__ == __ORDER_LITTLE_ENDIAN__
// Protobuf fixed32 is little-endian — direct load on LE platforms
memcpy(&val, ptr, 4);
#else
val = encode_uint32(ptr[3], ptr[2], ptr[1], ptr[0]);
#endif
if (!this->decode_32bit(field_id, Proto32Bit(val))) {
ESP_LOGV(TAG, "Cannot decode 32-bit field %" PRIu32 " with value %" PRIu32 "!", field_id, val);
}
ptr += 4;
break;
}
default:
ESP_LOGV(TAG, "Invalid field type %" PRIu32 " at offset %ld", field_type, (long) (ptr - buffer));
return;
}
this->decode_field(tag, data, scalar);
}
}
+166 -228
View File
@@ -170,43 +170,40 @@ class ProtoVarInt {
class ProtoMessage;
class ProtoSize;
/// Case label for decode_field(): the wire tag of a field, so a field that arrives with another wire
/// type matches no case.
constexpr uint32_t proto_tag(uint32_t field_id, uint32_t wire_type) { return (field_id << 3) | wire_type; }
/// One decoded field: the payload pointer and a scalar holding the varint or fixed32 value, or the
/// length of a length-delimited field. The wire type in the tag says which applies; accessors do not check.
class ProtoFieldValue {
class ProtoLengthDelimited {
public:
ProtoFieldValue(const uint8_t *data, proto_varint_value_t scalar) : data_(data), scalar_(scalar) {}
explicit ProtoLengthDelimited(const uint8_t *value, size_t length) : value_(value), length_(length) {}
std::string as_string() const { return std::string(reinterpret_cast<const char *>(this->value_), this->length_); }
proto_varint_value_t as_varint() const { return this->scalar_; }
// A bool is sent as 0 or 1, so the low word is enough and saves a second compare with 64 bit varints
bool as_bool() const { return static_cast<uint32_t>(this->scalar_) != 0; }
// Direct access to raw data without string allocation
const uint8_t *data() const { return this->value_; }
size_t size() const { return this->length_; }
// Length-delimited accessors
const uint8_t *data() const { return this->data_; }
size_t size() const { return static_cast<size_t>(this->scalar_); }
std::string as_string() const { return std::string(reinterpret_cast<const char *>(this->data_), this->size()); }
/// Decode the length-delimited payload into a message instance.
/// Decode the length-delimited data into a message instance.
/// Template preserves concrete type so decode() resolves statically.
template<typename T> void decode_to_message(T &msg) const { msg.decode(this->data_, this->size()); }
template<typename T> void decode_to_message(T &msg) const;
// Fixed32 accessors
uint32_t as_fixed32() const { return static_cast<uint32_t>(this->scalar_); }
int32_t as_sfixed32() const { return static_cast<int32_t>(this->as_fixed32()); }
protected:
const uint8_t *const value_;
const size_t length_;
};
class Proto32Bit {
public:
explicit Proto32Bit(uint32_t value) : value_(value) {}
uint32_t as_fixed32() const { return this->value_; }
int32_t as_sfixed32() const { return static_cast<int32_t>(this->value_); }
float as_float() const {
union {
uint32_t raw;
float value;
} s{};
s.raw = this->as_fixed32();
s.raw = this->value_;
return s.value;
}
private:
const uint8_t *data_;
proto_varint_value_t scalar_;
protected:
const uint32_t value_;
};
// NOTE: Proto64Bit class removed - wire type 1 (64-bit fixed) not supported
@@ -255,7 +252,7 @@ class ProtoWriteBuffer {
*
* Following https://protobuf.dev/programming-guides/encoding/#structure
*/
void encode_field_raw(uint32_t field_id, uint32_t type) { this->encode_varint_raw(proto_tag(field_id, type)); }
void encode_field_raw(uint32_t field_id, uint32_t type) { this->encode_varint_raw((field_id << 3) | type); }
/// Single-pass encode for repeated submessage elements.
/// Thin template wrapper; all buffer work is in the non-template core.
template<typename T> void encode_sub_message(uint32_t field_id, const T &value);
@@ -290,31 +287,19 @@ class ProtoWriteBuffer {
uint8_t *pos_;
};
// A four byte unaligned store is a memcpy call on ESP-IDF (-fno-builtin-memcpy) and on ARM cores without
// unaligned access (Cortex-M0+, ARM9), so those targets share one outlined byte store helper per fixed32
// field. Elsewhere the write inlines to a single store, or on ESP8266 to a few stores that measured
// faster than a call, so it stays inline.
#if defined(USE_ESP32) || (defined(__arm__) && !defined(__ARM_FEATURE_UNALIGNED))
#define PROTO_OUTLINE_FOR_SIZE __attribute__((noinline))
#define PROTO_FIXED32_BYTE_STORES true
#else
#define PROTO_OUTLINE_FOR_SIZE inline
#define PROTO_FIXED32_BYTE_STORES false
#endif
// Varint encoding thresholds — used by both proto_encode_* free functions and ProtoSize.
constexpr uint32_t VARINT_MAX_1_BYTE = 1 << 7; // 128
constexpr uint32_t VARINT_MAX_2_BYTE = 1 << 14; // 16384
/// Static encode helpers for the generated encode bodies. Each takes the write cursor by value and
/// returns it advanced, so outlined calls at -Os chain through the return register instead of a
/// stack slot. Helpers without a _force suffix skip fields holding the proto3 default.
/// Static encode helpers for generated encode() functions.
/// Generated code hoists buffer.pos_ into a local uint8_t *__restrict__ pos,
/// then calls these methods which take pos by reference. No struct, no overhead.
/// For sub-messages, pos is synced back to buffer before the call and reloaded after.
class ProtoEncode {
public:
/// Write a multi-byte varint directly through a pos pointer.
template<typename T>
[[nodiscard]] static inline uint8_t *encode_varint_raw_loop(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
T value) {
static inline void encode_varint_raw_loop(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, T value) {
do {
PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
*pos++ = static_cast<uint8_t>(value | 0x80);
@@ -322,49 +307,48 @@ class ProtoEncode {
} while (value > 0x7F);
PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
*pos++ = static_cast<uint8_t>(value);
return pos;
}
[[nodiscard]] static inline uint8_t *ESPHOME_ALWAYS_INLINE
encode_varint_raw(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint32_t value) {
static inline void ESPHOME_ALWAYS_INLINE encode_varint_raw(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t value) {
if (value < VARINT_MAX_1_BYTE) [[likely]] {
PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
*pos++ = static_cast<uint8_t>(value);
return pos;
return;
}
return encode_varint_raw_loop(pos PROTO_ENCODE_DEBUG_ARG, value);
encode_varint_raw_loop(pos PROTO_ENCODE_DEBUG_ARG, value);
}
/// Encode a varint that is expected to be 1-2 bytes (e.g. zigzag RSSI, small lengths).
[[nodiscard]] static inline uint8_t *ESPHOME_ALWAYS_INLINE
encode_varint_raw_short(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint32_t value) {
static inline void ESPHOME_ALWAYS_INLINE encode_varint_raw_short(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t value) {
if (value < VARINT_MAX_1_BYTE) [[likely]] {
PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
*pos++ = static_cast<uint8_t>(value);
return pos;
return;
}
if (value < VARINT_MAX_2_BYTE) [[likely]] {
PROTO_ENCODE_CHECK_BOUNDS(pos, 2);
*pos++ = static_cast<uint8_t>(value | 0x80);
*pos++ = static_cast<uint8_t>(value >> 7);
return pos;
return;
}
return encode_varint_raw_loop(pos PROTO_ENCODE_DEBUG_ARG, value);
encode_varint_raw_loop(pos PROTO_ENCODE_DEBUG_ARG, value);
}
[[nodiscard]] static inline uint8_t *ESPHOME_ALWAYS_INLINE
encode_varint_raw_64(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint64_t value) {
static inline void ESPHOME_ALWAYS_INLINE encode_varint_raw_64(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM,
uint64_t value) {
if (value < VARINT_MAX_1_BYTE) [[likely]] {
PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
*pos++ = static_cast<uint8_t>(value);
return pos;
return;
}
return encode_varint_raw_loop(pos PROTO_ENCODE_DEBUG_ARG, value);
encode_varint_raw_loop(pos PROTO_ENCODE_DEBUG_ARG, value);
}
/// Encode a 48-bit MAC address (stored in a uint64) as varint.
/// Real MAC addresses occupy the full 48 bits (OUI in upper 24), so the
/// fast path -- any non-zero bit in the top 6 of 48 -- emits exactly 7 bytes
/// with no per-byte branch. Falls back to the general loop otherwise.
/// Caller must guarantee value fits in 48 bits (checked in debug builds).
[[nodiscard]] static inline uint8_t *ESPHOME_ALWAYS_INLINE
encode_varint_raw_48bit(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint64_t value) {
static inline void ESPHOME_ALWAYS_INLINE encode_varint_raw_48bit(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM,
uint64_t value) {
#ifdef ESPHOME_DEBUG_API
assert(value < (1ULL << (MAC_ADDRESS_SIZE * 8)) && "encode_varint_raw_48bit: value exceeds 48 bits");
#endif
@@ -379,39 +363,38 @@ class ProtoEncode {
pos[4] = static_cast<uint8_t>((value >> 28) | 0x80);
pos[5] = static_cast<uint8_t>((value >> 35) | 0x80);
pos[6] = static_cast<uint8_t>(value >> 42);
return pos + 7;
pos += 7;
return;
}
return encode_varint_raw_64(pos PROTO_ENCODE_DEBUG_ARG, value);
encode_varint_raw_64(pos PROTO_ENCODE_DEBUG_ARG, value);
}
[[nodiscard]] static inline uint8_t *ESPHOME_ALWAYS_INLINE
encode_field_raw(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id, uint32_t type) {
return encode_varint_raw(pos PROTO_ENCODE_DEBUG_ARG, proto_tag(field_id, type));
static inline void ESPHOME_ALWAYS_INLINE encode_field_raw(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t field_id, uint32_t type) {
encode_varint_raw(pos PROTO_ENCODE_DEBUG_ARG, (field_id << 3) | type);
}
/// Write a single precomputed tag byte. Tag must be < 128.
[[nodiscard]] static inline uint8_t *ESPHOME_ALWAYS_INLINE
write_raw_byte(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint8_t b) {
static inline void ESPHOME_ALWAYS_INLINE write_raw_byte(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM,
uint8_t b) {
PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
*pos++ = b;
return pos;
}
/// Reserve one byte for later backpatch (e.g., sub-message length).
/// Advances pos past the reserved byte without writing a value.
[[nodiscard]] static inline uint8_t *ESPHOME_ALWAYS_INLINE
reserve_byte(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM) {
static inline void ESPHOME_ALWAYS_INLINE reserve_byte(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM) {
PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
return pos + 1;
pos++;
}
/// Write raw bytes to the buffer (no tag, no length prefix).
[[nodiscard]] static inline uint8_t *ESPHOME_ALWAYS_INLINE
encode_raw(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, const void *data, size_t len) {
static inline void ESPHOME_ALWAYS_INLINE encode_raw(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM,
const void *data, size_t len) {
PROTO_ENCODE_CHECK_BOUNDS(pos, len);
std::memcpy(pos, data, len);
return pos + len;
pos += len;
}
/// Encode tag + 1-byte length + raw string data. For strings with max_data_length < 128.
/// Tag must be a single-byte varint (< 128). Always encodes (no zero check).
[[nodiscard]] static inline uint8_t *encode_short_string_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint8_t tag, const StringRef &ref) {
static inline void encode_short_string_force(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint8_t tag,
const StringRef &ref) {
#ifdef ESPHOME_DEBUG_API
assert(ref.size() < 128 && "encode_short_string_force: string exceeds max_data_length < 128");
#endif
@@ -419,191 +402,137 @@ class ProtoEncode {
pos[0] = tag;
pos[1] = static_cast<uint8_t>(ref.size());
std::memcpy(pos + 2, ref.c_str(), ref.size());
return pos + 2 + ref.size();
pos += 2 + ref.size();
}
/// Write a precomputed tag byte + 32-bit value. Outlined on embedded: one copy beats inline stores per field.
[[nodiscard]] static PROTO_OUTLINE_FOR_SIZE uint8_t *write_tag_and_fixed32(
uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint8_t tag, uint32_t value) {
/// Write a precomputed tag byte + 32-bit value in one operation.
static inline void ESPHOME_ALWAYS_INLINE write_tag_and_fixed32(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM,
uint8_t tag, uint32_t value) {
PROTO_ENCODE_CHECK_BOUNDS(pos, 5);
pos[0] = tag;
write_fixed32_le(pos + 1, value);
return pos + 5;
#if __BYTE_ORDER__ == __ORDER_LITTLE_ENDIAN__
std::memcpy(pos + 1, &value, 4);
#else
pos[1] = static_cast<uint8_t>(value & 0xFF);
pos[2] = static_cast<uint8_t>((value >> 8) & 0xFF);
pos[3] = static_cast<uint8_t>((value >> 16) & 0xFF);
pos[4] = static_cast<uint8_t>((value >> 24) & 0xFF);
#endif
pos += 5;
}
[[nodiscard]] static inline uint8_t *encode_string_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t field_id, const char *string, size_t len) {
pos = encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 2); // type 2: Length-delimited string
static inline void encode_string(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
const char *string, size_t len, bool force = false) {
if (len == 0 && !force)
return;
encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 2); // type 2: Length-delimited string
// NOLINTNEXTLINE(readability-inconsistent-ifelse-braces) -- false positive on [[likely]] attribute
if (len < VARINT_MAX_1_BYTE) [[likely]] {
PROTO_ENCODE_CHECK_BOUNDS(pos, 1 + len);
*pos++ = static_cast<uint8_t>(len);
} else {
pos = encode_varint_raw_loop(pos PROTO_ENCODE_DEBUG_ARG, len);
encode_varint_raw_loop(pos PROTO_ENCODE_DEBUG_ARG, len);
PROTO_ENCODE_CHECK_BOUNDS(pos, len);
}
std::memcpy(pos, string, len);
return pos + len;
pos += len;
}
[[nodiscard]] static inline uint8_t *encode_string(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t field_id, const char *string, size_t len) {
if (len == 0)
return pos;
return encode_string_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, string, len);
static inline void encode_string(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
const std::string &value, bool force = false) {
encode_string(pos PROTO_ENCODE_DEBUG_ARG, field_id, value.data(), value.size(), force);
}
[[nodiscard]] static inline uint8_t *encode_string_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t field_id, const std::string &value) {
return encode_string_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, value.data(), value.size());
static inline void encode_string(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
const StringRef &ref, bool force = false) {
encode_string(pos PROTO_ENCODE_DEBUG_ARG, field_id, ref.c_str(), ref.size(), force);
}
[[nodiscard]] static inline uint8_t *encode_string(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t field_id, const StringRef &ref) {
return encode_string(pos PROTO_ENCODE_DEBUG_ARG, field_id, ref.c_str(), ref.size());
static inline void encode_bytes(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
const uint8_t *data, size_t len, bool force = false) {
encode_string(pos PROTO_ENCODE_DEBUG_ARG, field_id, reinterpret_cast<const char *>(data), len, force);
}
[[nodiscard]] static inline uint8_t *encode_string_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t field_id, const StringRef &ref) {
return encode_string_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, ref.c_str(), ref.size());
static inline void encode_uint32(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
uint32_t value, bool force = false) {
if (value == 0 && !force)
return;
encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 0);
encode_varint_raw(pos PROTO_ENCODE_DEBUG_ARG, value);
}
[[nodiscard]] static inline uint8_t *encode_bytes(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t field_id, const uint8_t *data, size_t len) {
return encode_string(pos PROTO_ENCODE_DEBUG_ARG, field_id, reinterpret_cast<const char *>(data), len);
static inline void encode_uint64(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
uint64_t value, bool force = false) {
if (value == 0 && !force)
return;
encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 0);
encode_varint_raw_64(pos PROTO_ENCODE_DEBUG_ARG, value);
}
[[nodiscard]] static inline uint8_t *encode_bytes_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t field_id, const uint8_t *data, size_t len) {
return encode_string_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, reinterpret_cast<const char *>(data), len);
}
[[nodiscard]] static inline uint8_t *encode_uint32_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t field_id, uint32_t value) {
pos = encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 0);
return encode_varint_raw(pos PROTO_ENCODE_DEBUG_ARG, value);
}
[[nodiscard]] static inline uint8_t *encode_uint32(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t field_id, uint32_t value) {
if (value == 0)
return pos;
return encode_uint32_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, value);
}
[[nodiscard]] static inline uint8_t *encode_uint64_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t field_id, uint64_t value) {
pos = encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 0);
return encode_varint_raw_64(pos PROTO_ENCODE_DEBUG_ARG, value);
}
[[nodiscard]] static inline uint8_t *encode_uint64(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t field_id, uint64_t value) {
if (value == 0)
return pos;
return encode_uint64_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, value);
}
[[nodiscard]] static inline uint8_t *encode_bool_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t field_id, bool value) {
pos = encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 0);
static inline void encode_bool(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id, bool value,
bool force = false) {
if (!value && !force)
return;
encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 0);
PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
*pos++ = value ? 0x01 : 0x00;
return pos;
}
[[nodiscard]] static inline uint8_t *encode_bool(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t field_id, bool value) {
if (!value)
return pos;
return encode_bool_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, value);
}
/// Tag + fixed32 for multi-byte tags; single-byte tags use write_tag_and_fixed32.
[[nodiscard]] static PROTO_OUTLINE_FOR_SIZE uint8_t *encode_fixed32_force(
uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id, uint32_t value) {
pos = encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 5);
static inline void encode_fixed32(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
uint32_t value, bool force = false) {
if (value == 0 && !force)
return;
encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 5);
PROTO_ENCODE_CHECK_BOUNDS(pos, 4);
write_fixed32_le(pos, value);
return pos + 4;
}
[[nodiscard]] static inline uint8_t *encode_fixed32(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t field_id, uint32_t value) {
if (value == 0)
return pos;
return encode_fixed32_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, value);
#if __BYTE_ORDER__ == __ORDER_LITTLE_ENDIAN__
std::memcpy(pos, &value, 4);
pos += 4;
#else
*pos++ = (value >> 0) & 0xFF;
*pos++ = (value >> 8) & 0xFF;
*pos++ = (value >> 16) & 0xFF;
*pos++ = (value >> 24) & 0xFF;
#endif
}
// NOTE: Wire type 1 (64-bit fixed: double, fixed64, sfixed64) is intentionally
// not supported to reduce overhead on embedded systems. All ESPHome devices are
// 32-bit microcontrollers where 64-bit operations are expensive. If 64-bit support
// is needed in the future, the necessary encoding/decoding functions must be added.
[[nodiscard]] static inline uint8_t *encode_float(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t field_id, float value) {
return encode_fixed32(pos PROTO_ENCODE_DEBUG_ARG, field_id, float_to_raw(value));
static inline void encode_float(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id, float value,
bool force = false) {
uint32_t raw = float_to_raw(value);
if (raw == 0 && !force)
return;
encode_fixed32(pos PROTO_ENCODE_DEBUG_ARG, field_id, raw);
}
[[nodiscard]] static inline uint8_t *encode_float_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t field_id, float value) {
return encode_fixed32_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, float_to_raw(value));
}
[[nodiscard]] static inline uint8_t *encode_int32_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t field_id, int32_t value) {
static inline void encode_int32(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id, int32_t value,
bool force = false) {
if (value < 0) {
// negative int32 is always 10 byte long
return encode_uint64_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, static_cast<uint64_t>(value));
encode_uint64(pos PROTO_ENCODE_DEBUG_ARG, field_id, static_cast<uint64_t>(value), force);
return;
}
return encode_uint32_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, static_cast<uint32_t>(value));
encode_uint32(pos PROTO_ENCODE_DEBUG_ARG, field_id, static_cast<uint32_t>(value), force);
}
[[nodiscard]] static inline uint8_t *encode_int32(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t field_id, int32_t value) {
if (value == 0)
return pos;
return encode_int32_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, value);
static inline void encode_int64(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id, int64_t value,
bool force = false) {
encode_uint64(pos PROTO_ENCODE_DEBUG_ARG, field_id, static_cast<uint64_t>(value), force);
}
[[nodiscard]] static inline uint8_t *encode_int64(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t field_id, int64_t value) {
return encode_uint64(pos PROTO_ENCODE_DEBUG_ARG, field_id, static_cast<uint64_t>(value));
static inline void encode_sint32(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
int32_t value, bool force = false) {
encode_uint32(pos PROTO_ENCODE_DEBUG_ARG, field_id, encode_zigzag32(value), force);
}
[[nodiscard]] static inline uint8_t *encode_int64_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t field_id, int64_t value) {
return encode_uint64_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, static_cast<uint64_t>(value));
static inline void encode_sint64(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
int64_t value, bool force = false) {
encode_uint64(pos PROTO_ENCODE_DEBUG_ARG, field_id, encode_zigzag64(value), force);
}
[[nodiscard]] static inline uint8_t *encode_sint32(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t field_id, int32_t value) {
return encode_uint32(pos PROTO_ENCODE_DEBUG_ARG, field_id, encode_zigzag32(value));
}
[[nodiscard]] static inline uint8_t *encode_sint32_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t field_id, int32_t value) {
return encode_uint32_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, encode_zigzag32(value));
}
[[nodiscard]] static inline uint8_t *encode_sint64(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t field_id, int64_t value) {
return encode_uint64(pos PROTO_ENCODE_DEBUG_ARG, field_id, encode_zigzag64(value));
}
[[nodiscard]] static inline uint8_t *encode_sint64_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t field_id, int64_t value) {
return encode_uint64_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, encode_zigzag64(value));
}
/// Sub-message encoding: sync pos to buffer, delegate, read the cursor back.
/// Sub-message encoding: sync pos to buffer, delegate, get pos from return value.
template<typename T>
[[nodiscard]] static inline uint8_t *encode_sub_message(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
ProtoWriteBuffer &buffer, uint32_t field_id, const T &value) {
static inline void encode_sub_message(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, ProtoWriteBuffer &buffer,
uint32_t field_id, const T &value) {
buffer.set_pos(pos);
buffer.encode_sub_message(field_id, value);
return buffer.get_pos();
pos = buffer.get_pos();
}
template<typename T>
[[nodiscard]] static inline uint8_t *encode_optional_sub_message(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
ProtoWriteBuffer &buffer, uint32_t field_id,
const T &value) {
static inline void encode_optional_sub_message(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM,
ProtoWriteBuffer &buffer, uint32_t field_id, const T &value) {
buffer.set_pos(pos);
buffer.encode_optional_sub_message(field_id, value);
return buffer.get_pos();
}
private:
/// Unaligned little endian store of four bytes: byte stores where the outlined helper lives (ESP-IDF, ARM
/// without unaligned access), otherwise a memcpy the compiler folds into one store. Callers bounds check
/// and advance the cursor themselves.
static inline void ESPHOME_ALWAYS_INLINE write_fixed32_le(uint8_t *__restrict__ pos, uint32_t value) {
if constexpr (PROTO_FIXED32_BYTE_STORES) {
// Spelled out so the outlined helper does not itself become a memcpy call
pos[0] = static_cast<uint8_t>(value);
pos[1] = static_cast<uint8_t>(value >> 8);
pos[2] = static_cast<uint8_t>(value >> 16);
pos[3] = static_cast<uint8_t>(value >> 24);
} else {
const uint32_t le = convert_little_endian(value);
__builtin_memcpy(pos, &le, 4);
}
pos = buffer.get_pos();
}
};
#undef PROTO_OUTLINE_FOR_SIZE
#undef PROTO_FIXED32_BYTE_STORES
#ifdef HAS_PROTO_MESSAGE_DUMP
/**
@@ -695,12 +624,11 @@ class DumpBuffer {
class ProtoMessage {
public:
// Non-virtual defaults for messages with no fields; generated classes hide all four. The
// static encode_msg/calc_size_msg take const void * so &T::encode_msg needs no thunk.
static uint8_t *encode_msg(const void *self, ProtoWriteBuffer &buffer PROTO_ENCODE_DEBUG_PARAM) {
return buffer.get_pos();
}
static uint32_t calc_size_msg(const void *self) { return 0; }
// Non-virtual defaults for messages with no fields.
// Concrete message classes hide these with their own implementations.
// All call sites use templates to preserve the concrete type, so virtual
// dispatch is not needed. This eliminates per-message vtable entries for
// encode/calculate_size, saving ~1.3 KB of flash across all message types.
uint8_t *encode(ProtoWriteBuffer &buffer PROTO_ENCODE_DEBUG_PARAM) const { return buffer.get_pos(); }
uint32_t calculate_size() const { return 0; }
#ifdef HAS_PROTO_MESSAGE_DUMP
@@ -735,10 +663,10 @@ class ProtoDecodableMessage : public ProtoMessage {
protected:
~ProtoDecodableMessage() = default;
/// Store one decoded field; \p scalar is the varint or fixed32 value, or the length of the
/// length-delimited payload at \p data. An unknown field or wrong wire type matches no case and is skipped.
/// Three register arguments keep the decode loop free of spills.
virtual void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) {}
virtual bool decode_varint(uint32_t field_id, proto_varint_value_t value) { return false; }
virtual bool decode_length(uint32_t field_id, ProtoLengthDelimited value) { return false; }
virtual bool decode_32bit(uint32_t field_id, Proto32Bit value) { return false; }
// NOTE: decode_64bit removed - wire type 1 not supported
};
class ProtoSize {
@@ -864,7 +792,7 @@ class ProtoSize {
* @return The number of bytes needed to encode the field ID and wire type
*/
static constexpr uint32_t field(uint32_t field_id, uint32_t type) {
uint32_t tag = proto_tag(field_id, type & WIRE_TYPE_MASK);
uint32_t tag = (field_id << 3) | (type & WIRE_TYPE_MASK);
return varint(tag);
}
@@ -948,14 +876,24 @@ class ProtoSize {
// Implementation of methods that depend on ProtoSize being fully defined
// Encode thunk — converts void* back to concrete type for direct encode() call
template<typename T> uint8_t *proto_encode_msg(const void *msg, ProtoWriteBuffer &buf PROTO_ENCODE_DEBUG_PARAM) {
return static_cast<const T *>(msg)->encode(buf PROTO_ENCODE_DEBUG_ARG);
}
// Thin template wrapper; delegates to non-template core in proto.cpp.
template<typename T> inline void ProtoWriteBuffer::encode_sub_message(uint32_t field_id, const T &value) {
this->encode_sub_message(field_id, &value, &T::encode_msg);
this->encode_sub_message(field_id, &value, &proto_encode_msg<T>);
}
// Thin template wrapper; delegates to non-template core.
template<typename T> inline void ProtoWriteBuffer::encode_optional_sub_message(uint32_t field_id, const T &value) {
this->encode_optional_sub_message(field_id, T::calc_size_msg(&value), &value, &T::encode_msg);
this->encode_optional_sub_message(field_id, value.calculate_size(), &value, &proto_encode_msg<T>);
}
// Template decode_to_message - preserves concrete type so decode() resolves statically
template<typename T> void ProtoLengthDelimited::decode_to_message(T &msg) const {
msg.decode(this->value_, this->length_);
}
template<typename T> const char *proto_enum_to_string(T value);
+206 -154
View File
@@ -9,6 +9,10 @@ namespace esphome::atm90e32 {
static const char *const TAG = "atm90e32";
static const LogString *offset_calibration_name(bool power_offsets) {
return power_offsets ? LOG_STR("Power offset") : LOG_STR("Offset");
}
static uint32_t pref_hash(const char *prefix, const char *name_space) {
auto hash = fnv1_hash(prefix);
return fnv1_hash_extend(hash, name_space);
@@ -203,13 +207,12 @@ void ATM90E32Component::setup() {
// Initialize flash storage for power offset calibrations
uint32_t po_hash = pref_hash("_power_offset_calibration_", cs);
this->power_offset_pref_ = global_preferences->make_preference<PowerOffsetCalibration[3]>(po_hash, true);
this->power_offset_pref_ = global_preferences->make_preference<OffsetCalibration[3]>(po_hash, true);
bool migrated_power_offset = false;
if (has_distinct_legacy_namespace) {
uint32_t legacy_po_hash = pref_hash("_power_offset_calibration_", legacy_cs);
auto legacy_power_offset_pref =
global_preferences->make_preference<PowerOffsetCalibration[3]>(legacy_po_hash, true);
PowerOffsetCalibration power_offset_data[3]{};
auto legacy_power_offset_pref = global_preferences->make_preference<OffsetCalibration[3]>(legacy_po_hash, true);
OffsetCalibration power_offset_data[3]{};
int migration_status =
migrate_legacy_pref_if_needed(this->power_offset_pref_, legacy_power_offset_pref, &power_offset_data);
migrated_power_offset = migration_status > 0;
@@ -224,20 +227,20 @@ void ATM90E32Component::setup() {
global_preferences->sync();
}
this->restore_offset_calibrations_();
this->restore_power_offset_calibrations_();
this->restore_offset_calibrations_(OffsetCalibrationType::OFFSET_CALIBRATION_TYPE_VOLTAGE_CURRENT);
this->restore_offset_calibrations_(OffsetCalibrationType::OFFSET_CALIBRATION_TYPE_POWER);
} else {
ESP_LOGI(TAG, "[CALIBRATION][%s] Power & Voltage/Current offset calibration is disabled. Using config file values.",
cs);
for (uint8_t phase = 0; phase < 3; ++phase) {
this->write16_(this->voltage_offset_registers[phase],
static_cast<uint16_t>(this->offset_phase_[phase].voltage_offset_));
static_cast<uint16_t>(this->offset_phase_[phase].first_offset));
this->write16_(this->current_offset_registers[phase],
static_cast<uint16_t>(this->offset_phase_[phase].current_offset_));
static_cast<uint16_t>(this->offset_phase_[phase].second_offset));
this->write16_(this->power_offset_registers[phase],
static_cast<uint16_t>(this->power_offset_phase_[phase].active_power_offset));
static_cast<uint16_t>(this->power_offset_phase_[phase].first_offset));
this->write16_(this->reactive_power_offset_registers[phase],
static_cast<uint16_t>(this->power_offset_phase_[phase].reactive_power_offset));
static_cast<uint16_t>(this->power_offset_phase_[phase].second_offset));
}
}
@@ -317,8 +320,8 @@ void ATM90E32Component::log_calibration_status_() {
cs);
for (uint8_t phase = 0; phase < 3; ++phase) {
ESP_LOGW(TAG, "[CALIBRATION][%s] | %c | %6d | %6d | %6d | %6d |", cs, 'A' + phase,
this->config_offset_phase_[phase].voltage_offset_, this->offset_phase_[phase].voltage_offset_,
this->config_offset_phase_[phase].current_offset_, this->offset_phase_[phase].current_offset_);
this->config_offset_phase_[phase].first_offset, this->offset_phase_[phase].first_offset,
this->config_offset_phase_[phase].second_offset, this->offset_phase_[phase].second_offset);
}
ESP_LOGW(TAG,
"[CALIBRATION][%s] ===============================================================================", cs);
@@ -335,10 +338,8 @@ void ATM90E32Component::log_calibration_status_() {
cs);
for (uint8_t phase = 0; phase < 3; ++phase) {
ESP_LOGW(TAG, "[CALIBRATION][%s] | %c | %6d | %6d | %6d | %6d |", cs, 'A' + phase,
this->config_power_offset_phase_[phase].active_power_offset,
this->power_offset_phase_[phase].active_power_offset,
this->config_power_offset_phase_[phase].reactive_power_offset,
this->power_offset_phase_[phase].reactive_power_offset);
this->config_power_offset_phase_[phase].first_offset, this->power_offset_phase_[phase].first_offset,
this->config_power_offset_phase_[phase].second_offset, this->power_offset_phase_[phase].second_offset);
}
ESP_LOGW(TAG,
"[CALIBRATION][%s] ===============================================================================", cs);
@@ -372,7 +373,7 @@ void ATM90E32Component::log_calibration_status_() {
ESP_LOGI(TAG, "[CALIBRATION][%s] --------------------------------------------------------------", cs);
for (uint8_t phase = 0; phase < 3; phase++) {
ESP_LOGI(TAG, "[CALIBRATION][%s] | %c | %6d | %6d |", cs, 'A' + phase,
this->offset_phase_[phase].voltage_offset_, this->offset_phase_[phase].current_offset_);
this->offset_phase_[phase].first_offset, this->offset_phase_[phase].second_offset);
}
ESP_LOGI(TAG, "[CALIBRATION][%s] ==============================================================\\n", cs);
}
@@ -385,8 +386,7 @@ void ATM90E32Component::log_calibration_status_() {
ESP_LOGI(TAG, "[CALIBRATION][%s] ---------------------------------------------------------------------", cs);
for (uint8_t phase = 0; phase < 3; phase++) {
ESP_LOGI(TAG, "[CALIBRATION][%s] | %c | %6d | %6d |", cs, 'A' + phase,
this->power_offset_phase_[phase].active_power_offset,
this->power_offset_phase_[phase].reactive_power_offset);
this->power_offset_phase_[phase].first_offset, this->power_offset_phase_[phase].second_offset);
}
ESP_LOGI(TAG, "[CALIBRATION][%s] =====================================================================\n", cs);
}
@@ -756,36 +756,68 @@ void ATM90E32Component::save_gain_calibration_to_memory_() {
}
}
void ATM90E32Component::save_offset_calibration_to_memory_() {
void ATM90E32Component::finish_offset_calibration_(const OffsetCalibration (&previous)[3], bool previous_restored,
bool previous_using_saved, OffsetCalibrationType type) {
const bool power_offsets = type == OffsetCalibrationType::OFFSET_CALIBRATION_TYPE_POWER;
const char *cs = this->get_calibration_id_();
bool success = this->offset_pref_.save(&this->offset_phase_);
global_preferences->sync();
if (success) {
this->using_saved_calibrations_ = true;
this->restored_offset_calibration_ = true;
for (bool &phase : this->offset_calibration_mismatch_)
phase = false;
ESP_LOGI(TAG, "[CALIBRATION][%s] Offset calibration saved to memory.", cs);
} else {
this->using_saved_calibrations_ = false;
ESP_LOGE(TAG, "[CALIBRATION][%s] Failed to save offset calibration to memory!", cs);
}
}
const LogString *name = offset_calibration_name(power_offsets);
OffsetCalibration(*offsets)[3] = power_offsets ? &this->power_offset_phase_ : &this->offset_phase_;
ESPPreferenceObject *preference = power_offsets ? &this->power_offset_pref_ : &this->offset_pref_;
bool *has_stored =
power_offsets ? &this->has_stored_power_offset_calibration_ : &this->has_stored_offset_calibration_;
bool *restored = power_offsets ? &this->restored_power_offset_calibration_ : &this->restored_offset_calibration_;
bool *mismatches = power_offsets ? this->power_offset_calibration_mismatch_ : this->offset_calibration_mismatch_;
void ATM90E32Component::save_power_offset_calibration_to_memory_() {
const char *cs = this->get_calibration_id_();
bool success = this->power_offset_pref_.save(&this->power_offset_phase_);
global_preferences->sync();
if (success) {
this->using_saved_calibrations_ = true;
this->restored_power_offset_calibration_ = true;
for (bool &phase : this->power_offset_calibration_mismatch_)
phase = false;
ESP_LOGI(TAG, "[CALIBRATION][%s] Power offset calibration saved to memory.", cs);
} else {
this->using_saved_calibrations_ = false;
ESP_LOGE(TAG, "[CALIBRATION][%s] Failed to save power offset calibration to memory!", cs);
const bool writes_verified = this->verify_offset_writes_(type);
bool saved = false;
bool synced = false;
if (writes_verified) {
saved = preference->save(offsets);
synced = global_preferences->sync();
}
if (writes_verified && saved && synced) {
this->using_saved_calibrations_ = true;
*has_stored = true;
*restored = true;
for (uint8_t phase = 0; phase < 3; phase++)
mismatches[phase] = false;
ESP_LOGI(TAG, "[CALIBRATION][%s] %s calibration saved to memory. %s calibration completed and verified.", cs,
LOG_STR_ARG(name), LOG_STR_ARG(name));
return;
}
if (writes_verified) {
ESP_LOGE(TAG, "[CALIBRATION][%s] Failed to save %s calibration to memory!", cs, LOG_STR_ARG(name));
}
for (uint8_t phase = 0; phase < 3; phase++) {
this->write_offsets_to_registers_(phase, previous[phase].first_offset, previous[phase].second_offset, type);
}
const bool rollback_verified = this->verify_offset_writes_(type);
bool rollback_persisted = false;
if (writes_verified) {
OffsetCalibration rollback[3]{};
prepare_offset_rollback(previous, previous_restored, rollback);
const bool rollback_saved = preference->save(&rollback);
const bool rollback_synced = global_preferences->sync();
rollback_persisted = rollback_saved && rollback_synced;
if (!rollback_saved || !rollback_synced) {
ESP_LOGE(TAG, "[CALIBRATION][%s] Failed to persist restored %s calibration values!", cs, LOG_STR_ARG(name));
}
}
*restored = previous_restored;
if (rollback_persisted)
*has_stored = previous_restored;
this->using_saved_calibrations_ = previous_using_saved;
if (!rollback_verified) {
ESP_LOGE(TAG, "[CALIBRATION][%s] %s calibration failed; rollback readback verification failed.", cs,
LOG_STR_ARG(name));
return;
}
ESP_LOGE(TAG, "[CALIBRATION][%s] %s calibration failed; previous values restored.", cs, LOG_STR_ARG(name));
}
void ATM90E32Component::run_offset_calibrations() {
@@ -803,11 +835,16 @@ void ATM90E32Component::run_offset_calibrations() {
ESP_LOGI(TAG, "[CALIBRATION][%s] | Phase | offset_voltage | offset_current |", cs);
ESP_LOGI(TAG, "[CALIBRATION][%s] ------------------------------------------------------------------", cs);
OffsetCalibration previous_offsets[3] = {this->offset_phase_[0], this->offset_phase_[1], this->offset_phase_[2]};
const bool previous_restored = this->restored_offset_calibration_;
const bool previous_using_saved = this->using_saved_calibrations_;
for (uint8_t phase = 0; phase < 3; phase++) {
int16_t voltage_offset = calibrate_offset(phase, true);
int16_t current_offset = calibrate_offset(phase, false);
this->write_offsets_to_registers_(phase, voltage_offset, current_offset);
this->write_offsets_to_registers_(phase, voltage_offset, current_offset,
OffsetCalibrationType::OFFSET_CALIBRATION_TYPE_VOLTAGE_CURRENT);
ESP_LOGI(TAG, "[CALIBRATION][%s] | %c | %6d | %6d |", cs, 'A' + phase, voltage_offset,
current_offset);
@@ -815,7 +852,8 @@ void ATM90E32Component::run_offset_calibrations() {
ESP_LOGI(TAG, "[CALIBRATION][%s] ==================================================================\n", cs);
this->save_offset_calibration_to_memory_();
this->finish_offset_calibration_(previous_offsets, previous_restored, previous_using_saved,
OffsetCalibrationType::OFFSET_CALIBRATION_TYPE_VOLTAGE_CURRENT);
}
void ATM90E32Component::run_power_offset_calibrations() {
@@ -834,18 +872,25 @@ void ATM90E32Component::run_power_offset_calibrations() {
ESP_LOGI(TAG, "[CALIBRATION][%s] | Phase | offset_active_power | offset_reactive_power |", cs);
ESP_LOGI(TAG, "[CALIBRATION][%s] ---------------------------------------------------------------------", cs);
OffsetCalibration previous_offsets[3] = {this->power_offset_phase_[0], this->power_offset_phase_[1],
this->power_offset_phase_[2]};
const bool previous_restored = this->restored_power_offset_calibration_;
const bool previous_using_saved = this->using_saved_calibrations_;
for (uint8_t phase = 0; phase < 3; ++phase) {
int16_t active_offset = calibrate_power_offset(phase, false);
int16_t reactive_offset = calibrate_power_offset(phase, true);
this->write_power_offsets_to_registers_(phase, active_offset, reactive_offset);
this->write_offsets_to_registers_(phase, active_offset, reactive_offset,
OffsetCalibrationType::OFFSET_CALIBRATION_TYPE_POWER);
ESP_LOGI(TAG, "[CALIBRATION][%s] | %c | %6d | %6d |", cs, 'A' + phase, active_offset,
reactive_offset);
}
ESP_LOGI(TAG, "[CALIBRATION][%s] =====================================================================\n", cs);
this->save_power_offset_calibration_to_memory_();
this->finish_offset_calibration_(previous_offsets, previous_restored, previous_using_saved,
OffsetCalibrationType::OFFSET_CALIBRATION_TYPE_POWER);
}
void ATM90E32Component::write_gains_to_registers_() {
@@ -859,35 +904,26 @@ void ATM90E32Component::write_gains_to_registers_() {
this->write16_(ATM90E32_REGISTER_CFGREGACCEN, 0x0000);
}
void ATM90E32Component::write_offsets_to_registers_(uint8_t phase, int16_t voltage_offset, int16_t current_offset) {
// Save to runtime
this->offset_phase_[phase].voltage_offset_ = voltage_offset;
this->phase_[phase].voltage_offset_ = voltage_offset;
void ATM90E32Component::write_offsets_to_registers_(uint8_t phase, int16_t first_offset, int16_t second_offset,
OffsetCalibrationType type) {
const bool power_offsets = type == OffsetCalibrationType::OFFSET_CALIBRATION_TYPE_POWER;
OffsetCalibration &offsets = power_offsets ? this->power_offset_phase_[phase] : this->offset_phase_[phase];
offsets.first_offset = first_offset;
offsets.second_offset = second_offset;
if (power_offsets) {
this->phase_[phase].active_power_offset_ = first_offset;
this->phase_[phase].reactive_power_offset_ = second_offset;
} else {
this->phase_[phase].voltage_offset_ = first_offset;
this->phase_[phase].current_offset_ = second_offset;
}
// Save to flash-storable struct
this->offset_phase_[phase].current_offset_ = current_offset;
this->phase_[phase].current_offset_ = current_offset;
// Write to registers
const uint16_t *first_registers = power_offsets ? this->power_offset_registers : this->voltage_offset_registers;
const uint16_t *second_registers =
power_offsets ? this->reactive_power_offset_registers : this->current_offset_registers;
this->write16_(ATM90E32_REGISTER_CFGREGACCEN, 0x55AA);
this->write16_(voltage_offset_registers[phase], static_cast<uint16_t>(voltage_offset));
this->write16_(current_offset_registers[phase], static_cast<uint16_t>(current_offset));
this->write16_(ATM90E32_REGISTER_CFGREGACCEN, 0x0000);
}
void ATM90E32Component::write_power_offsets_to_registers_(uint8_t phase, int16_t p_offset, int16_t q_offset) {
// Save to runtime
this->phase_[phase].active_power_offset_ = p_offset;
this->phase_[phase].reactive_power_offset_ = q_offset;
// Save to flash-storable struct
this->power_offset_phase_[phase].active_power_offset = p_offset;
this->power_offset_phase_[phase].reactive_power_offset = q_offset;
// Write to registers
this->write16_(ATM90E32_REGISTER_CFGREGACCEN, 0x55AA);
this->write16_(this->power_offset_registers[phase], static_cast<uint16_t>(p_offset));
this->write16_(this->reactive_power_offset_registers[phase], static_cast<uint16_t>(q_offset));
this->write16_(first_registers[phase], static_cast<uint16_t>(first_offset));
this->write16_(second_registers[phase], static_cast<uint16_t>(second_offset));
this->write16_(ATM90E32_REGISTER_CFGREGACCEN, 0x0000);
}
@@ -947,89 +983,78 @@ void ATM90E32Component::restore_gain_calibrations_() {
ESP_LOGW(TAG, "[CALIBRATION][%s] No stored gain calibrations found. Using config file values.", cs);
}
void ATM90E32Component::restore_offset_calibrations_() {
void ATM90E32Component::restore_offset_calibrations_(OffsetCalibrationType type) {
const bool power_offsets = type == OffsetCalibrationType::OFFSET_CALIBRATION_TYPE_POWER;
const char *cs = this->get_calibration_id_();
const LogString *name = power_offsets ? LOG_STR("power offset") : LOG_STR("offset");
OffsetCalibration(*offsets)[3] = power_offsets ? &this->power_offset_phase_ : &this->offset_phase_;
OffsetCalibration(*config_offsets)[3] =
power_offsets ? &this->config_power_offset_phase_ : &this->config_offset_phase_;
ESPPreferenceObject *preference = power_offsets ? &this->power_offset_pref_ : &this->offset_pref_;
bool *has_stored =
power_offsets ? &this->has_stored_power_offset_calibration_ : &this->has_stored_offset_calibration_;
bool *restored = power_offsets ? &this->restored_power_offset_calibration_ : &this->restored_offset_calibration_;
bool *mismatches = power_offsets ? this->power_offset_calibration_mismatch_ : this->offset_calibration_mismatch_;
const bool *has_first = power_offsets ? this->has_config_active_power_offset_ : this->has_config_voltage_offset_;
const bool *has_second = power_offsets ? this->has_config_reactive_power_offset_ : this->has_config_current_offset_;
for (uint8_t i = 0; i < 3; ++i)
this->config_offset_phase_[i] = this->offset_phase_[i];
bool have_data = this->offset_pref_.load(&this->offset_phase_);
(*config_offsets)[i] = (*offsets)[i];
const bool have_data = preference->load(offsets);
bool all_zero = true;
if (have_data) {
for (auto &phase : this->offset_phase_) {
if (phase.voltage_offset_ != 0 || phase.current_offset_ != 0) {
for (const auto &phase : *offsets) {
if (phase.first_offset != 0 || phase.second_offset != 0) {
all_zero = false;
break;
}
}
}
if (have_data && !all_zero) {
this->restored_offset_calibration_ = true;
for (uint8_t phase = 0; phase < 3; phase++) {
auto &offset = this->offset_phase_[phase];
bool mismatch = false;
if (this->has_config_voltage_offset_[phase] &&
offset.voltage_offset_ != this->config_offset_phase_[phase].voltage_offset_)
mismatch = true;
if (this->has_config_current_offset_[phase] &&
offset.current_offset_ != this->config_offset_phase_[phase].current_offset_)
mismatch = true;
if (mismatch)
this->offset_calibration_mismatch_[phase] = true;
*has_stored = have_data && !all_zero;
*restored = false;
for (uint8_t phase = 0; phase < 3; phase++) {
mismatches[phase] = false;
if (*has_stored) {
mismatches[phase] =
(has_first[phase] && (*offsets)[phase].first_offset != (*config_offsets)[phase].first_offset) ||
(has_second[phase] && (*offsets)[phase].second_offset != (*config_offsets)[phase].second_offset);
}
} else {
}
if (!*has_stored) {
for (uint8_t phase = 0; phase < 3; phase++)
this->offset_phase_[phase] = this->config_offset_phase_[phase];
ESP_LOGW(TAG, "[CALIBRATION][%s] No stored offset calibrations found. Using default values.", cs);
(*offsets)[phase] = (*config_offsets)[phase];
ESP_LOGW(TAG, "[CALIBRATION][%s] No stored %s calibrations found. Using default values.", cs, LOG_STR_ARG(name));
}
for (uint8_t phase = 0; phase < 3; phase++) {
write_offsets_to_registers_(phase, this->offset_phase_[phase].voltage_offset_,
this->offset_phase_[phase].current_offset_);
this->write_offsets_to_registers_(phase, (*offsets)[phase].first_offset, (*offsets)[phase].second_offset, type);
}
}
void ATM90E32Component::restore_power_offset_calibrations_() {
const char *cs = this->get_calibration_id_();
for (uint8_t i = 0; i < 3; ++i)
this->config_power_offset_phase_[i] = this->power_offset_phase_[i];
bool have_data = this->power_offset_pref_.load(&this->power_offset_phase_);
bool all_zero = true;
if (have_data) {
for (auto &phase : this->power_offset_phase_) {
if (phase.active_power_offset != 0 || phase.reactive_power_offset != 0) {
all_zero = false;
break;
}
}
const bool initial_values_verified = this->verify_offset_writes_(type);
if (initial_values_verified) {
const auto state = resolve_offset_restore_state(*has_stored, true, false);
*restored = state.restored;
ESP_LOGI(TAG, "[CALIBRATION][%s] %s calibration values verified.", cs, LOG_STR_ARG(name));
return;
}
if (have_data && !all_zero) {
this->restored_power_offset_calibration_ = true;
for (uint8_t phase = 0; phase < 3; ++phase) {
auto &offset = this->power_offset_phase_[phase];
bool mismatch = false;
if (this->has_config_active_power_offset_[phase] &&
offset.active_power_offset != this->config_power_offset_phase_[phase].active_power_offset)
mismatch = true;
if (this->has_config_reactive_power_offset_[phase] &&
offset.reactive_power_offset != this->config_power_offset_phase_[phase].reactive_power_offset)
mismatch = true;
if (mismatch)
this->power_offset_calibration_mismatch_[phase] = true;
}
this->using_saved_calibrations_ = false;
for (uint8_t phase = 0; phase < 3; phase++)
mismatches[phase] = false;
for (uint8_t phase = 0; phase < 3; phase++) {
(*offsets)[phase] = (*config_offsets)[phase];
this->write_offsets_to_registers_(phase, (*offsets)[phase].first_offset, (*offsets)[phase].second_offset, type);
}
const auto state = resolve_offset_restore_state(*has_stored, false, this->verify_offset_writes_(type));
*restored = state.restored;
if (state.values_verified) {
ESP_LOGE(TAG, "[CALIBRATION][%s] %s calibration restore failed verification; config values verified.", cs,
LOG_STR_ARG(name));
} else {
for (uint8_t phase = 0; phase < 3; ++phase)
this->power_offset_phase_[phase] = this->config_power_offset_phase_[phase];
ESP_LOGW(TAG, "[CALIBRATION][%s] No stored power offsets found. Using default values.", cs);
}
for (uint8_t phase = 0; phase < 3; ++phase) {
write_power_offsets_to_registers_(phase, this->power_offset_phase_[phase].active_power_offset,
this->power_offset_phase_[phase].reactive_power_offset);
ESP_LOGE(TAG, "[CALIBRATION][%s] %s calibration restore and config fallback both failed verification.", cs,
LOG_STR_ARG(name));
}
}
@@ -1084,14 +1109,14 @@ void ATM90E32Component::clear_gain_calibrations() {
void ATM90E32Component::clear_offset_calibrations() {
const char *cs = this->get_calibration_id_();
if (!this->restored_offset_calibration_) {
if (!this->has_stored_offset_calibration_) {
ESP_LOGI(TAG, "[CALIBRATION][%s] No stored offset calibrations to clear. Current values:", cs);
ESP_LOGI(TAG, "[CALIBRATION][%s] --------------------------------------------------------------", cs);
ESP_LOGI(TAG, "[CALIBRATION][%s] | Phase | offset_voltage | offset_current |", cs);
ESP_LOGI(TAG, "[CALIBRATION][%s] --------------------------------------------------------------", cs);
for (uint8_t phase = 0; phase < 3; phase++) {
ESP_LOGI(TAG, "[CALIBRATION][%s] | %c | %6d | %6d |", cs, 'A' + phase,
this->offset_phase_[phase].voltage_offset_, this->offset_phase_[phase].current_offset_);
this->offset_phase_[phase].first_offset, this->offset_phase_[phase].second_offset);
}
ESP_LOGI(TAG, "[CALIBRATION][%s] ==============================================================\n", cs);
return;
@@ -1104,10 +1129,11 @@ void ATM90E32Component::clear_offset_calibrations() {
for (uint8_t phase = 0; phase < 3; phase++) {
int16_t voltage_offset =
this->has_config_voltage_offset_[phase] ? this->config_offset_phase_[phase].voltage_offset_ : 0;
this->has_config_voltage_offset_[phase] ? this->config_offset_phase_[phase].first_offset : 0;
int16_t current_offset =
this->has_config_current_offset_[phase] ? this->config_offset_phase_[phase].current_offset_ : 0;
this->write_offsets_to_registers_(phase, voltage_offset, current_offset);
this->has_config_current_offset_[phase] ? this->config_offset_phase_[phase].second_offset : 0;
this->write_offsets_to_registers_(phase, voltage_offset, current_offset,
OffsetCalibrationType::OFFSET_CALIBRATION_TYPE_VOLTAGE_CURRENT);
ESP_LOGI(TAG, "[CALIBRATION][%s] | %c | %6d | %6d |", cs, 'A' + phase, voltage_offset,
current_offset);
}
@@ -1117,6 +1143,7 @@ void ATM90E32Component::clear_offset_calibrations() {
this->offset_pref_.save(&zero_offsets); // Clear stored values in flash
global_preferences->sync();
this->has_stored_offset_calibration_ = false;
this->restored_offset_calibration_ = false;
for (bool &phase : this->offset_calibration_mismatch_)
phase = false;
@@ -1126,15 +1153,14 @@ void ATM90E32Component::clear_offset_calibrations() {
void ATM90E32Component::clear_power_offset_calibrations() {
const char *cs = this->get_calibration_id_();
if (!this->restored_power_offset_calibration_) {
if (!this->has_stored_power_offset_calibration_) {
ESP_LOGI(TAG, "[CALIBRATION][%s] No stored power offsets to clear. Current values:", cs);
ESP_LOGI(TAG, "[CALIBRATION][%s] ---------------------------------------------------------------------", cs);
ESP_LOGI(TAG, "[CALIBRATION][%s] | Phase | offset_active_power | offset_reactive_power |", cs);
ESP_LOGI(TAG, "[CALIBRATION][%s] ---------------------------------------------------------------------", cs);
for (uint8_t phase = 0; phase < 3; phase++) {
ESP_LOGI(TAG, "[CALIBRATION][%s] | %c | %6d | %6d |", cs, 'A' + phase,
this->power_offset_phase_[phase].active_power_offset,
this->power_offset_phase_[phase].reactive_power_offset);
this->power_offset_phase_[phase].first_offset, this->power_offset_phase_[phase].second_offset);
}
ESP_LOGI(TAG, "[CALIBRATION][%s] =====================================================================\n", cs);
return;
@@ -1147,20 +1173,21 @@ void ATM90E32Component::clear_power_offset_calibrations() {
for (uint8_t phase = 0; phase < 3; phase++) {
int16_t active_offset =
this->has_config_active_power_offset_[phase] ? this->config_power_offset_phase_[phase].active_power_offset : 0;
int16_t reactive_offset = this->has_config_reactive_power_offset_[phase]
? this->config_power_offset_phase_[phase].reactive_power_offset
: 0;
this->write_power_offsets_to_registers_(phase, active_offset, reactive_offset);
this->has_config_active_power_offset_[phase] ? this->config_power_offset_phase_[phase].first_offset : 0;
int16_t reactive_offset =
this->has_config_reactive_power_offset_[phase] ? this->config_power_offset_phase_[phase].second_offset : 0;
this->write_offsets_to_registers_(phase, active_offset, reactive_offset,
OffsetCalibrationType::OFFSET_CALIBRATION_TYPE_POWER);
ESP_LOGI(TAG, "[CALIBRATION][%s] | %c | %6d | %6d |", cs, 'A' + phase, active_offset,
reactive_offset);
}
ESP_LOGI(TAG, "[CALIBRATION][%s] =====================================================================\n", cs);
PowerOffsetCalibration zero_power_offsets[3]{{0, 0}, {0, 0}, {0, 0}};
OffsetCalibration zero_power_offsets[3]{{0, 0}, {0, 0}, {0, 0}};
this->power_offset_pref_.save(&zero_power_offsets);
global_preferences->sync();
this->has_stored_power_offset_calibration_ = false;
this->restored_power_offset_calibration_ = false;
for (bool &phase : this->power_offset_calibration_mismatch_)
phase = false;
@@ -1215,6 +1242,31 @@ bool ATM90E32Component::verify_gain_writes_() {
return success; // Return true if all writes were successful, false otherwise
}
bool ATM90E32Component::verify_offset_writes_(OffsetCalibrationType type) {
const bool power_offsets = type == OffsetCalibrationType::OFFSET_CALIBRATION_TYPE_POWER;
const char *cs = this->get_calibration_id_();
const LogString *name = offset_calibration_name(power_offsets);
const LogString *first_name = power_offsets ? LOG_STR("active") : LOG_STR("voltage");
const LogString *second_name = power_offsets ? LOG_STR("reactive") : LOG_STR("current");
const OffsetCalibration *offsets = power_offsets ? this->power_offset_phase_ : this->offset_phase_;
const uint16_t *first_registers = power_offsets ? this->power_offset_registers : this->voltage_offset_registers;
const uint16_t *second_registers =
power_offsets ? this->reactive_power_offset_registers : this->current_offset_registers;
bool success = true;
for (uint8_t phase = 0; phase < 3; phase++) {
const uint16_t first = this->read16_(first_registers[phase]);
const uint16_t second = this->read16_(second_registers[phase]);
if (!offset_register_value_matches(first, offsets[phase].first_offset) ||
!offset_register_value_matches(second, offsets[phase].second_offset)) {
ESP_LOGE(TAG, "[CALIBRATION][%s] %s readback failed for Phase %s: %s %d/%d, %s %d/%d.", cs, LOG_STR_ARG(name),
phase_labels[phase], LOG_STR_ARG(first_name), static_cast<int16_t>(first), offsets[phase].first_offset,
LOG_STR_ARG(second_name), static_cast<int16_t>(second), offsets[phase].second_offset);
success = false;
}
}
return success;
}
#ifdef USE_TEXT_SENSOR
void ATM90E32Component::check_phase_status() {
uint16_t state0 = this->read16_(ATM90E32_REGISTER_EMMSTATE0);
+49 -22
View File
@@ -13,6 +13,40 @@
namespace esphome::atm90e32 {
inline bool offset_register_value_matches(uint16_t actual, int16_t expected) {
return actual == static_cast<uint16_t>(expected);
}
struct OffsetCalibration {
int16_t first_offset{0};
int16_t second_offset{0};
};
static_assert(sizeof(OffsetCalibration[3]) == 12, "Offset calibration preference layout must remain compatible");
enum class OffsetCalibrationType : uint8_t {
OFFSET_CALIBRATION_TYPE_VOLTAGE_CURRENT,
OFFSET_CALIBRATION_TYPE_POWER,
};
struct OffsetRestoreState {
bool restored;
bool values_verified;
};
inline OffsetRestoreState resolve_offset_restore_state(bool has_stored_values, bool initial_values_verified,
bool fallback_values_verified) {
if (initial_values_verified)
return {has_stored_values, true};
return {false, fallback_values_verified};
}
inline void prepare_offset_rollback(const OffsetCalibration (&previous)[3], bool had_stored_values,
OffsetCalibration (&rollback)[3]) {
for (uint8_t phase = 0; phase < 3; phase++)
rollback[phase] = had_stored_values ? previous[phase] : OffsetCalibration{};
}
class ATM90E32Component final : public PollingComponent,
public spi::SPIDevice<spi::BIT_ORDER_MSB_FIRST, spi::CLOCK_POLARITY_HIGH,
spi::CLOCK_PHASE_TRAILING, spi::DATA_RATE_1MHZ> {
@@ -71,19 +105,19 @@ class ATM90E32Component final : public PollingComponent,
this->has_config_current_gain_[phase] = true;
}
void set_voltage_offset(uint8_t phase, int16_t offset) {
this->offset_phase_[phase].voltage_offset_ = offset;
this->offset_phase_[phase].first_offset = offset;
this->has_config_voltage_offset_[phase] = true;
}
void set_current_offset(uint8_t phase, int16_t offset) {
this->offset_phase_[phase].current_offset_ = offset;
this->offset_phase_[phase].second_offset = offset;
this->has_config_current_offset_[phase] = true;
}
void set_active_power_offset(uint8_t phase, int16_t offset) {
this->power_offset_phase_[phase].active_power_offset = offset;
this->power_offset_phase_[phase].first_offset = offset;
this->has_config_active_power_offset_[phase] = true;
}
void set_reactive_power_offset(uint8_t phase, int16_t offset) {
this->power_offset_phase_[phase].reactive_power_offset = offset;
this->power_offset_phase_[phase].second_offset = offset;
this->has_config_reactive_power_offset_[phase] = true;
}
void set_freq_sensor(sensor::Sensor *freq_sensor) { freq_sensor_ = freq_sensor; }
@@ -171,16 +205,16 @@ class ATM90E32Component final : public PollingComponent,
float get_chip_temperature_();
bool get_publish_interval_flag_() { return publish_interval_flag_; };
void set_publish_interval_flag_(bool flag) { publish_interval_flag_ = flag; };
void restore_offset_calibrations_();
void restore_power_offset_calibrations_();
void restore_offset_calibrations_(OffsetCalibrationType type);
void restore_gain_calibrations_();
void save_offset_calibration_to_memory_();
void save_gain_calibration_to_memory_();
void save_power_offset_calibration_to_memory_();
void write_offsets_to_registers_(uint8_t phase, int16_t voltage_offset, int16_t current_offset);
void write_power_offsets_to_registers_(uint8_t phase, int16_t p_offset, int16_t q_offset);
void finish_offset_calibration_(const OffsetCalibration (&previous)[3], bool previous_restored,
bool previous_using_saved, OffsetCalibrationType type);
void write_offsets_to_registers_(uint8_t phase, int16_t first_offset, int16_t second_offset,
OffsetCalibrationType type);
void write_gains_to_registers_();
bool verify_gain_writes_();
bool verify_offset_writes_(OffsetCalibrationType type);
bool validate_spi_read_(uint16_t expected, const char *context = nullptr);
void log_calibration_status_();
const char *get_calibration_id_();
@@ -219,19 +253,10 @@ class ATM90E32Component final : public PollingComponent,
uint32_t cumulative_reverse_active_energy_{0};
} phase_[3];
struct OffsetCalibration {
int16_t voltage_offset_{0};
int16_t current_offset_{0};
} offset_phase_[3];
OffsetCalibration offset_phase_[3];
OffsetCalibration config_offset_phase_[3];
struct PowerOffsetCalibration {
int16_t active_power_offset{0};
int16_t reactive_power_offset{0};
} power_offset_phase_[3];
PowerOffsetCalibration config_power_offset_phase_[3];
OffsetCalibration power_offset_phase_[3];
OffsetCalibration config_power_offset_phase_[3];
struct GainCalibration {
uint16_t voltage_gain{1};
@@ -265,6 +290,8 @@ class ATM90E32Component final : public PollingComponent,
bool enable_offset_calibration_{false};
bool enable_gain_calibration_{false};
const char *instance_id_{nullptr};
bool has_stored_offset_calibration_{false};
bool has_stored_power_offset_calibration_{false};
bool restored_offset_calibration_{false};
bool restored_power_offset_calibration_{false};
bool restored_gain_calibration_{false};
+4 -3
View File
@@ -313,9 +313,10 @@ FileDecoderState AudioDecoder::decode_mp3_() {
this->output_transfer_buffer_->increase_buffer_length(
this->audio_stream_info_.value().frames_to_bytes(samples_decoded));
}
} else if (result == micro_mp3::MP3_STREAM_INFO_READY) {
// First successful header parse: capture stream info and resize the output buffer to fit one full frame.
// microMP3 always outputs 16-bit PCM.
} else if (result == micro_mp3::MP3_STREAM_INFO_READY || result == micro_mp3::MP3_STREAM_INFO_CHANGED) {
// Header parsed: capture stream info and resize the output buffer to fit one full frame.
// microMP3 always outputs 16-bit PCM. MP3_STREAM_INFO_CHANGED is handled identically: despite its
// negative value it is documented as recoverable, so it must not reach the catch-all below.
this->audio_stream_info_ =
audio::AudioStreamInfo(16, this->mp3_decoder_->get_channels(), this->mp3_decoder_->get_sample_rate());
this->free_buffer_required_ =
+24 -10
View File
@@ -22,6 +22,23 @@ class Automation {
static const char *const TAG;
};
// Base for nodes that never read the parent's services.
// The parent releases its services only once every node reports Established, so a node that never
// reports it keeps that memory allocated for the life of the connection.
class BLEClientServicelessNode : public BLEClientNode {
public:
// Final so that Established is always reported on SEARCH_CMPL, before the derived node sees the event.
void gattc_event_handler(esp_gattc_cb_event_t event, esp_gatt_if_t gattc_if, esp_ble_gattc_cb_param_t *param) final {
if (event == ESP_GATTC_SEARCH_CMPL_EVT)
this->node_state = espbt::ClientState::ESTABLISHED;
this->on_gattc_event(event, gattc_if, param);
}
protected:
// Derived nodes handle GATT events here rather than by overriding the handler above.
virtual void on_gattc_event(esp_gattc_cb_event_t event, esp_gatt_if_t gattc_if, esp_ble_gattc_cb_param_t *param) {}
};
// implement on_connect automation.
class BLEClientConnectTrigger final : public Trigger<>, public BLEClientNode {
public:
@@ -61,7 +78,7 @@ class BLEClientDisconnectTrigger final : public Trigger<>, public BLEClientNode
}
};
class BLEClientPasskeyRequestTrigger final : public Trigger<>, public BLEClientNode {
class BLEClientPasskeyRequestTrigger final : public Trigger<>, public BLEClientServicelessNode {
public:
explicit BLEClientPasskeyRequestTrigger(BLEClient *parent) { parent->register_ble_node(this); }
void loop() override {}
@@ -71,7 +88,7 @@ class BLEClientPasskeyRequestTrigger final : public Trigger<>, public BLEClientN
}
};
class BLEClientPasskeyNotificationTrigger final : public Trigger<uint32_t>, public BLEClientNode {
class BLEClientPasskeyNotificationTrigger final : public Trigger<uint32_t>, public BLEClientServicelessNode {
public:
explicit BLEClientPasskeyNotificationTrigger(BLEClient *parent) { parent->register_ble_node(this); }
void loop() override {}
@@ -82,7 +99,7 @@ class BLEClientPasskeyNotificationTrigger final : public Trigger<uint32_t>, publ
}
};
class BLEClientNumericComparisonRequestTrigger final : public Trigger<uint32_t>, public BLEClientNode {
class BLEClientNumericComparisonRequestTrigger final : public Trigger<uint32_t>, public BLEClientServicelessNode {
public:
explicit BLEClientNumericComparisonRequestTrigger(BLEClient *parent) { parent->register_ble_node(this); }
void loop() override {}
@@ -315,19 +332,17 @@ template<typename... Ts> class BLEClientRemoveBondAction final : public Action<T
BLEClient *parent_{nullptr};
};
template<typename... Ts> class BLEClientConnectAction final : public Action<Ts...>, public BLEClientNode {
template<typename... Ts> class BLEClientConnectAction final : public Action<Ts...>, public BLEClientServicelessNode {
public:
BLEClientConnectAction(BLEClient *ble_client) {
ble_client->register_ble_node(this);
ble_client_ = ble_client;
}
void gattc_event_handler(esp_gattc_cb_event_t event, esp_gatt_if_t gattc_if,
esp_ble_gattc_cb_param_t *param) override {
void on_gattc_event(esp_gattc_cb_event_t event, esp_gatt_if_t gattc_if, esp_ble_gattc_cb_param_t *param) override {
if (this->num_running_ == 0)
return;
switch (event) {
case ESP_GATTC_SEARCH_CMPL_EVT:
this->node_state = espbt::ClientState::ESTABLISHED;
this->parent()->run_later([this]() { this->play_next_tuple_(this->var_); });
break;
// if the connection is closed, terminate the automation chain.
@@ -364,14 +379,13 @@ template<typename... Ts> class BLEClientConnectAction final : public Action<Ts..
std::tuple<Ts...> var_{};
};
template<typename... Ts> class BLEClientDisconnectAction final : public Action<Ts...>, public BLEClientNode {
template<typename... Ts> class BLEClientDisconnectAction final : public Action<Ts...>, public BLEClientServicelessNode {
public:
BLEClientDisconnectAction(BLEClient *ble_client) {
ble_client->register_ble_node(this);
ble_client_ = ble_client;
}
void gattc_event_handler(esp_gattc_cb_event_t event, esp_gatt_if_t gattc_if,
esp_ble_gattc_cb_param_t *param) override {
void on_gattc_event(esp_gattc_cb_event_t event, esp_gatt_if_t gattc_if, esp_ble_gattc_cb_param_t *param) override {
if (this->num_running_ == 0)
return;
switch (event) {
@@ -6,6 +6,7 @@ namespace esphome::dallas_temp {
static const char *const TAG = "dallas.temp.sensor";
static const uint8_t DALLAS_MODEL_DS18S20 = 0x10;
static const uint8_t DALLAS_MODEL_DS18B20 = 0x28;
static const uint8_t DALLAS_COMMAND_START_CONVERSION = 0x44;
static const uint8_t DALLAS_COMMAND_READ_SCRATCH_PAD = 0xBE;
static const uint8_t DALLAS_COMMAND_WRITE_SCRATCH_PAD = 0x4E;
@@ -154,7 +155,14 @@ float DallasTemperatureSensor::get_temp_c_() {
default:
break;
}
// undocumented test for powerup measurement of 85
// https://github.com/cpetrich/counterfeit_DS18B20#solution-to-the-85-c-problem
if ((this->address_ & 0xff) == DALLAS_MODEL_DS18B20) {
if ((temp == 85 * 16) && (this->scratch_pad_[6] == 0xc)) {
ESP_LOGD(TAG, "dropping reading caused by sensor reset");
return NAN;
}
}
return temp / 16.0f;
}
+6 -2
View File
@@ -66,11 +66,15 @@ const char *DebugComponent::get_reset_reason_(std::span<char, RESET_REASON_BUFFE
unsigned reason = esp_reset_reason();
if (reason < sizeof(RESET_REASONS) / sizeof(RESET_REASONS[0])) {
if (reason == ESP_RST_SW) {
if (reason == ESP_RST_SW || reason == ESP_RST_WDT) {
// On some ESP32-S3 configurations (e.g. SPIRAM with fetch-instructions/rodata),
// esp_restart() intermittently produces RTCWDT_RTC_RST (ESP_RST_WDT) instead of
// ESP_RST_SW. Check the stored reboot source for both reset reasons so a software
// reboot that ends up as WDT still reports the correct source.
auto pref = global_preferences->make_preference(REBOOT_MAX_LEN,
fnv1_hash_extend(fnv1_hash(REBOOT_KEY), App.get_name().c_str()));
char reboot_source[REBOOT_MAX_LEN]{};
if (pref.load(&reboot_source)) {
if (pref.load(&reboot_source) && reboot_source[0] != '\0') {
reboot_source[REBOOT_MAX_LEN - 1] = '\0';
snprintf(buf, size, "Reboot request from %s", reboot_source);
} else {
+1 -5
View File
@@ -23,11 +23,7 @@ from esphome.const import (
)
from esphome.types import ConfigType
from . import ( # noqa: F401 pylint: disable=unused-import
CONF_DEBUG_ID,
FILTER_SOURCE_FILES,
DebugComponent,
)
from . import CONF_DEBUG_ID, FILTER_SOURCE_FILES, DebugComponent # noqa: F401 pylint: disable=unused-import
DEPENDENCIES = ["debug"]
+1 -5
View File
@@ -9,11 +9,7 @@ from esphome.const import (
)
from esphome.types import ConfigType
from . import ( # noqa: F401 pylint: disable=unused-import
CONF_DEBUG_ID,
FILTER_SOURCE_FILES,
DebugComponent,
)
from . import CONF_DEBUG_ID, FILTER_SOURCE_FILES, DebugComponent # noqa: F401 pylint: disable=unused-import
DEPENDENCIES = ["debug"]
+2 -9
View File
@@ -3,18 +3,11 @@ import esphome.codegen as cg
# Re-exported for the many esp32-side users; defined in esphome.const
# and esphome.espidf so the upload/logs fast path can use them without
# importing this package.
from esphome.const import ( # noqa: F401 # pylint: disable=unused-import
KEY_ESP32,
KEY_FLASH_SIZE,
KEY_IDF_VERSION,
KEY_VARIANT,
)
from esphome.const import KEY_ESP32, KEY_FLASH_SIZE, KEY_IDF_VERSION, KEY_VARIANT # noqa: F401 # pylint: disable=unused-import
# Back compat for external components only; in-tree callers import it
# from esphome.espidf directly.
from esphome.espidf import ( # noqa: F401 # pylint: disable=unused-import
variant_to_idf_target,
)
from esphome.espidf import variant_to_idf_target # noqa: F401 # pylint: disable=unused-import
KEY_BOARD = "board"
KEY_SDKCONFIG_OPTIONS = "sdkconfig_options"
@@ -91,7 +91,14 @@ void I2SAudioSpeakerBase::loop() {
this->speaker_task_handle_ = nullptr;
this->stop_i2s_driver_();
xEventGroupClearBits(this->event_group_, SpeakerEventGroupBits::ALL_BITS);
// ALL_BITS includes COMMAND_START. Take the bits from the clear itself, not from the snapshot at
// the top of loop(): the audio source's task can raise a start at any point above, including
// during stop_i2s_driver_(), and nothing would ever re-issue it.
const EventBits_t bits_before_clear = xEventGroupClearBits(this->event_group_, SpeakerEventGroupBits::ALL_BITS);
if (bits_before_clear & SpeakerEventGroupBits::COMMAND_START) {
ESP_LOGD(TAG, "Start requested while stopping; keeping the request");
xEventGroupSetBits(this->event_group_, SpeakerEventGroupBits::COMMAND_START);
}
this->status_clear_error();
this->on_task_stopped();
@@ -5,6 +5,7 @@
#include <esp_log.h>
#include <driver/uart.h>
#include <soc/soc_caps.h>
#ifdef USE_LOGGER_UART_SELECTION_USB_SERIAL_JTAG
#include <driver/usb_serial_jtag.h>
@@ -76,7 +77,11 @@ void init_uart(uart_port_t uart_num, uint32_t baud_rate, int tx_buffer_size) {
uart_config.parity = UART_PARITY_DISABLE;
uart_config.stop_bits = UART_STOP_BITS_1;
uart_config.flow_ctrl = UART_HW_FLOWCTRL_DISABLE;
#if SOC_UART_SUPPORT_XTAL_CLK
uart_config.source_clk = UART_SCLK_XTAL;
#else
uart_config.source_clk = UART_SCLK_DEFAULT;
#endif
uart_param_config(uart_num, &uart_config);
// The logger only writes to UART, never reads, so use the minimum RX buffer.
// ESP-IDF requires rx_buffer_size > UART_HW_FIFO_LEN (128 bytes).
+2 -1
View File
@@ -15,6 +15,7 @@ from ..defines import (
from ..types import LvCompound, LvType
from . import Widget, WidgetType, get_widgets
from .buttonmatrix import CONF_BUTTONMATRIX
from .label import CONF_LABEL
from .textarea import CONF_TEXTAREA, lv_textarea_t
CONF_KEYBOARD = "keyboard"
@@ -49,7 +50,7 @@ class KeyboardType(WidgetType):
)
def get_uses(self):
return CONF_KEYBOARD, CONF_TEXTAREA, CONF_BUTTONMATRIX
return CONF_KEYBOARD, CONF_TEXTAREA, CONF_BUTTONMATRIX, CONF_LABEL
async def to_code(self, w: Widget, config: dict):
add_lv_use("KEY_LISTENER")
+2 -1
View File
@@ -10,6 +10,7 @@ from ..types import lv_obj_t
from . import Widget, WidgetType
from .canvas import CONF_CANVAS
from .img import CONF_IMAGE
from .label import CONF_LABEL
CONF_QRCODE = "qrcode"
CONF_DARK_COLOR = "dark_color"
@@ -41,7 +42,7 @@ class QrCodeType(WidgetType):
)
def get_uses(self):
return CONF_CANVAS, CONF_IMAGE
return CONF_CANVAS, CONF_IMAGE, CONF_LABEL
async def to_code(self, w: Widget, config):
await w.set_property(
+2 -1
View File
@@ -28,6 +28,7 @@ from ..types import LV_EVENT, LvType, ObjUpdateAction, lv_obj_t, lv_obj_t_ptr
from . import Widget, WidgetType, add_widgets, get_widgets, set_obj_properties
from .button import button_spec
from .buttonmatrix import CONF_BUTTONMATRIX, buttonmatrix_spec
from .label import CONF_LABEL
from .obj import obj_spec
CONF_TABVIEW = "tabview"
@@ -74,7 +75,7 @@ class TabviewType(WidgetType):
)
def get_uses(self):
return CONF_BUTTONMATRIX, TYPE_FLEX, CONF_BUTTON
return CONF_BUTTONMATRIX, TYPE_FLEX, CONF_BUTTON, CONF_LABEL
async def to_code(self, w: Widget, config: dict):
await w.set_property(
+3
View File
@@ -67,6 +67,9 @@ void MQTTJSONLightComponent::send_discovery(JsonObject root, mqtt::SendDiscovery
if (traits.supports_color_mode(ColorMode::RGB_COLD_WARM_WHITE))
color_modes.add(ESPHOME_F("rgbww"));
if (traits.supports_color_capability(ColorCapability::BRIGHTNESS))
root[ESPHOME_F("brightness")] = true;
if (traits.supports_color_mode(ColorMode::COLOR_TEMPERATURE) ||
traits.supports_color_mode(ColorMode::COLD_WARM_WHITE)) {
root[MQTT_MIN_MIREDS] = traits.get_min_mireds();
+1 -6
View File
@@ -14,12 +14,7 @@ from esphome.const import (
)
from esphome.core import CORE, TimePeriod
from . import ( # noqa: F401 pylint: disable=unused-import
FILTER_SOURCE_FILES,
Nextion,
nextion_ns,
nextion_ref,
)
from . import FILTER_SOURCE_FILES, Nextion, nextion_ns, nextion_ref # noqa: F401 pylint: disable=unused-import
from .base_component import (
CONF_AUTO_WAKE_ON_TOUCH,
CONF_COMMAND_SPACING,
+9 -1
View File
@@ -6,6 +6,8 @@
#include <cstdint>
#include "esphome/core/log.h"
#include "noise_resume.h"
namespace esphome::noise {
using psk_t = std::array<uint8_t, 32>;
@@ -26,13 +28,19 @@ class NoiseContext {
/// psk points at 32 bytes that outlive the context (PROGMEM or caller owned
/// RAM); nullptr means no key. Runtime callers map the all-zeros key to
/// nullptr themselves; validation keeps it out of yaml.
void set_psk(const uint8_t *psk) { this->psk_ = psk; }
void set_psk(const uint8_t *psk) {
this->psk_ = psk;
// Resume tickets were minted under the old key; forget them
this->resume_cache_.clear();
}
/// Copy the key out (flash-aware on ESP8266); all zeros when none is set.
void load_psk(psk_t &out) const;
bool has_psk() const { return this->psk_ != nullptr; }
ResumeTicketCache &resume_cache() { return this->resume_cache_; }
protected:
const uint8_t *psk_{nullptr};
ResumeTicketCache resume_cache_;
};
/// Convert a noise error code to a readable error
+127
View File
@@ -0,0 +1,127 @@
#include "noise_resume.h"
#ifdef USE_NOISE
#include <cstring>
#include <noise/protocol.h>
#include "esphome/core/hal.h"
#include "esphome/core/helpers.h"
namespace esphome::noise {
const char RESUME_LABEL_OFFER[6] PROGMEM = "offer";
const char RESUME_LABEL_CONFIRM[8] PROGMEM = "confirm";
const char RESUME_LABEL_KEYS[5] PROGMEM = "keys";
bool resume_kdf(const uint8_t *secret, const char *label, size_t label_len, const uint8_t *a, size_t a_len,
const uint8_t *b, size_t b_len, const uint8_t *hash_in, size_t hash_in_len, uint8_t *out1,
size_t out1_len, uint8_t *out2) {
uint8_t data[RESUME_KDF_MAX_DATA];
uint8_t scratch[32];
size_t len = label_len + a_len + b_len;
progmem_memcpy(data, label, label_len);
std::memcpy(data + label_len, a, a_len);
std::memcpy(data + label_len + a_len, b, b_len);
NoiseHashState *hash = nullptr;
if (noise_hashstate_new_by_id(&hash, NOISE_HASH_SHA256) != NOISE_ERROR_NONE) {
return false;
}
int err = NOISE_ERROR_NONE;
if (hash_in != nullptr) {
err = noise_hashstate_hash_one(hash, hash_in, hash_in_len, data + len, 32);
len += 32;
}
if (err == NOISE_ERROR_NONE) {
err = noise_hashstate_hkdf(hash, secret, RESUME_SECRET_SIZE, data, len, out1, out1_len,
out2 != nullptr ? out2 : scratch, 32);
}
noise_hashstate_free(hash);
noise_clean(data, sizeof(data));
noise_clean(scratch, sizeof(scratch));
return err == NOISE_ERROR_NONE;
}
bool ResumeTicketCache::issue(ResumeTicket &out) {
if (!random_bytes(reinterpret_cast<uint8_t *>(&out), sizeof(out))) {
return false;
}
uint8_t slot = this->next_;
this->next_ = static_cast<uint8_t>((slot + 1) % SLOTS);
this->slots_[slot] = out;
this->used_mask_ |= static_cast<uint8_t>(1u << slot);
return true;
}
size_t ResumeTicketCache::try_accept(const uint8_t *offer, size_t offer_len, const uint8_t *prologue,
size_t prologue_len, uint8_t *out_ext, size_t out_capacity,
NoiseCipherState *&send_cipher, NoiseCipherState *&recv_cipher) {
if (offer_len != RESUME_OFFER_SIZE || offer[0] != RESUME_OFFER_VERSION || out_capacity < RESUME_ACCEPT_SIZE) {
return 0;
}
const uint8_t *session_id = offer + RESUME_OFFER_SESSION_ID_OFFSET;
const uint8_t *client_nonce = offer + RESUME_OFFER_NONCE_OFFSET;
ResumeTicket *ticket = nullptr;
for (uint8_t i = 0; i < SLOTS; i++) {
if ((this->used_mask_ & (1u << i)) &&
std::memcmp(this->slots_[i].session_id, session_id, RESUME_SESSION_ID_SIZE) == 0) {
ticket = &this->slots_[i];
this->used_mask_ &= static_cast<uint8_t>(~(1u << i));
break;
}
}
if (ticket == nullptr) {
return 0;
}
uint8_t expected[RESUME_MAC_SIZE];
bool ok = resume_compute_offer_mac(ticket->secret, session_id, client_nonce, expected) &&
noise_is_equal(expected, offer + RESUME_OFFER_MAC_OFFSET, RESUME_MAC_SIZE);
noise_clean(expected, sizeof(expected));
if (!ok) {
// Bad MAC: keep the ticket so a forger cannot burn it
this->used_mask_ |= static_cast<uint8_t>(1u << static_cast<uint8_t>(ticket - this->slots_));
return 0;
}
// The ticket is spent from here; any later failure falls back to the full
// handshake and the client gets a fresh one.
uint8_t *server_nonce = out_ext + 1;
uint8_t k_c2d[32];
uint8_t k_d2c[32];
out_ext[0] = RESUME_ACCEPT_VERSION;
ok = random_bytes(server_nonce, RESUME_NONCE_SIZE) &&
resume_compute_confirm_mac(ticket->secret, client_nonce, server_nonce, out_ext + 1 + RESUME_NONCE_SIZE) &&
resume_derive_keys(ticket->secret, client_nonce, server_nonce, prologue, prologue_len, k_c2d, k_d2c);
noise_clean(ticket, sizeof(*ticket));
if (ok) {
recv_cipher = resume_make_cipher(k_c2d);
send_cipher = resume_make_cipher(k_d2c);
ok = recv_cipher != nullptr && send_cipher != nullptr;
if (!ok) {
noise_cipherstate_free(recv_cipher);
noise_cipherstate_free(send_cipher);
recv_cipher = nullptr;
send_cipher = nullptr;
}
}
noise_clean(k_c2d, sizeof(k_c2d));
noise_clean(k_d2c, sizeof(k_d2c));
return ok ? RESUME_ACCEPT_SIZE : 0;
}
void ResumeTicketCache::clear() {
noise_clean(this->slots_, sizeof(this->slots_));
this->used_mask_ = 0;
}
NoiseCipherState *resume_make_cipher(const uint8_t *key) {
NoiseCipherState *cipher = nullptr;
if (noise_cipherstate_new_by_id(&cipher, NOISE_CIPHER_CHACHAPOLY) != NOISE_ERROR_NONE) {
return nullptr;
}
if (noise_cipherstate_init_key(cipher, key, 32) != NOISE_ERROR_NONE) {
noise_cipherstate_free(cipher);
return nullptr;
}
return cipher;
}
} // namespace esphome::noise
#endif // USE_NOISE
+132
View File
@@ -0,0 +1,132 @@
#pragma once
#include "esphome/core/defines.h"
#ifdef USE_NOISE
#include <cstddef>
#include <cstdint>
// Forward declaration matching <noise/protocol/cipherstate.h>; keeps noise-c
// headers out of everything that includes noise.h.
extern "C" {
typedef struct NoiseCipherState_s NoiseCipherState; // NOLINT(modernize-use-using)
}
namespace esphome::noise {
/** Session resume for the noise transports.
*
* After a full handshake the responder issues a single-use ticket over the
* encrypted channel. A client presents it in its next ClientHello and both
* sides derive the transport keys with HKDF-SHA256 alone, skipping the two
* curve25519 operations. Old peers ignore the extension bytes on both
* sides, so every mismatch degrades to a normal full handshake.
*
* HKDF is the Noise construction (noise_hashstate_hkdf). Derivations:
* offer_mac = HKDF(secret, "offer" || session_id || client_nonce).out1[:16]
* confirm_mac = HKDF(secret, "confirm" || client_nonce || server_nonce).out1[:16]
* k_c2d, k_d2c = HKDF(secret, "keys" || client_nonce || server_nonce || SHA256(prologue))
*
* An offering client sends handshake message 1 only after a decline.
* Resumed sessions have no ephemeral DH; the ticket is wiped on use.
*/
static constexpr uint8_t RESUME_OFFER_VERSION = 0x01;
static constexpr uint8_t RESUME_ACCEPT_VERSION = 0x01;
static constexpr size_t RESUME_SESSION_ID_SIZE = 8;
static constexpr size_t RESUME_NONCE_SIZE = 16;
static constexpr size_t RESUME_MAC_SIZE = 16;
static constexpr size_t RESUME_SECRET_SIZE = 32;
// ClientHello body: version | session_id | client_nonce | offer_mac
static constexpr size_t RESUME_OFFER_SIZE = 1 + RESUME_SESSION_ID_SIZE + RESUME_NONCE_SIZE + RESUME_MAC_SIZE; // 41
static constexpr size_t RESUME_OFFER_SESSION_ID_OFFSET = 1;
static constexpr size_t RESUME_OFFER_NONCE_OFFSET = RESUME_OFFER_SESSION_ID_OFFSET + RESUME_SESSION_ID_SIZE;
static constexpr size_t RESUME_OFFER_MAC_OFFSET = RESUME_OFFER_NONCE_OFFSET + RESUME_NONCE_SIZE;
// ServerHello trailing extension: version | server_nonce | confirm_mac
static constexpr size_t RESUME_ACCEPT_SIZE = 1 + RESUME_NONCE_SIZE + RESUME_MAC_SIZE; // 33
struct ResumeTicket {
uint8_t session_id[RESUME_SESSION_ID_SIZE];
uint8_t secret[RESUME_SECRET_SIZE];
};
// Sent on the wire as one blob: session_id || secret
static_assert(sizeof(ResumeTicket) == RESUME_SESSION_ID_SIZE + RESUME_SECRET_SIZE, "ticket must be packed");
/// Fixed-slot RAM cache of single-use resume tickets. Lost on reboot by
/// design: clients fall back to a full handshake.
class ResumeTicketCache {
public:
/// Generate a fresh ticket into out and store it, evicting the oldest
/// slot. Returns false (and stores nothing) if the RNG fails.
bool issue(ResumeTicket &out);
/// Accept a resume offer: verify and consume the ticket (single use; a
/// forged MAC never burns one), build both transport ciphers, and write
/// the ServerHello accept extension into out_ext. Returns the extension
/// length, or 0 (nothing allocated) on any miss, failure, or when
/// out_capacity is too small. Secrets are wiped internally.
size_t try_accept(const uint8_t *offer, size_t offer_len, const uint8_t *prologue, size_t prologue_len,
uint8_t *out_ext, size_t out_capacity, NoiseCipherState *&send_cipher,
NoiseCipherState *&recv_cipher);
/// Forget every ticket (PSK change).
void clear();
// Round robin; more clients than slots thrash and fall back to full handshakes
static constexpr uint8_t SLOTS = 2;
static_assert(SLOTS <= 8, "used_mask_ is uint8_t");
protected:
ResumeTicket slots_[SLOTS];
uint8_t used_mask_{0};
uint8_t next_{0};
};
/// HKDF labels, PROGMEM on ESP8266.
extern const char RESUME_LABEL_OFFER[6];
extern const char RESUME_LABEL_CONFIRM[8];
extern const char RESUME_LABEL_KEYS[5];
// Largest KDF input: "keys" || client_nonce || server_nonce || SHA256(prologue)
static constexpr size_t RESUME_KDF_MAX_DATA =
sizeof(RESUME_LABEL_KEYS) - 1 + RESUME_NONCE_SIZE + RESUME_NONCE_SIZE + 32;
/// Noise-construction HKDF-SHA256 keyed with the ticket secret over
/// label || a || b [|| SHA256(hash_in)], at most RESUME_KDF_MAX_DATA. out2 == nullptr means MAC only.
bool resume_kdf(const uint8_t *secret, const char *label, size_t label_len, const uint8_t *a, size_t a_len,
const uint8_t *b, size_t b_len, const uint8_t *hash_in, size_t hash_in_len, uint8_t *out1,
size_t out1_len, uint8_t *out2);
/// offer_mac for the ClientHello resume offer (what a client computes and
/// try_accept checks).
inline bool resume_compute_offer_mac(const uint8_t *secret, const uint8_t *session_id, const uint8_t *client_nonce,
uint8_t *out_mac) {
static_assert(sizeof(RESUME_LABEL_OFFER) - 1 + RESUME_SESSION_ID_SIZE + RESUME_NONCE_SIZE <= RESUME_KDF_MAX_DATA,
"KDF buffer");
return resume_kdf(secret, RESUME_LABEL_OFFER, sizeof(RESUME_LABEL_OFFER) - 1, session_id, RESUME_SESSION_ID_SIZE,
client_nonce, RESUME_NONCE_SIZE, nullptr, 0, out_mac, RESUME_MAC_SIZE, nullptr);
}
/// confirm_mac for the ServerHello extension.
inline bool resume_compute_confirm_mac(const uint8_t *secret, const uint8_t *client_nonce, const uint8_t *server_nonce,
uint8_t *out_mac) {
static_assert(sizeof(RESUME_LABEL_CONFIRM) - 1 + RESUME_NONCE_SIZE + RESUME_NONCE_SIZE <= RESUME_KDF_MAX_DATA,
"KDF buffer");
return resume_kdf(secret, RESUME_LABEL_CONFIRM, sizeof(RESUME_LABEL_CONFIRM) - 1, client_nonce, RESUME_NONCE_SIZE,
server_nonce, RESUME_NONCE_SIZE, nullptr, 0, out_mac, RESUME_MAC_SIZE, nullptr);
}
/// Derive the transport keys. k_c2d encrypts client-to-device traffic,
/// k_d2c device-to-client.
inline bool resume_derive_keys(const uint8_t *secret, const uint8_t *client_nonce, const uint8_t *server_nonce,
const uint8_t *prologue, size_t prologue_len, uint8_t *k_c2d, uint8_t *k_d2c) {
static_assert(sizeof(RESUME_LABEL_KEYS) - 1 + RESUME_NONCE_SIZE + RESUME_NONCE_SIZE + 32 <= RESUME_KDF_MAX_DATA,
"KDF buffer");
return resume_kdf(secret, RESUME_LABEL_KEYS, sizeof(RESUME_LABEL_KEYS) - 1, client_nonce, RESUME_NONCE_SIZE,
server_nonce, RESUME_NONCE_SIZE, prologue, prologue_len, k_c2d, 32, k_d2c);
}
/// Build a ChaChaPoly cipher state keyed with key (32 bytes); nullptr on
/// failure. Nonce counter starts at 0, exactly like a post-split cipher.
NoiseCipherState *resume_make_cipher(const uint8_t *key);
} // namespace esphome::noise
#endif // USE_NOISE
+88 -21
View File
@@ -18,6 +18,16 @@ void RFBridgeComponent::ack_() {
}
bool RFBridgeComponent::parse_bridge_byte_(uint8_t byte) {
if (this->bucket_frame_candidate_ && byte == RF_CODE_START) {
// A queued next frame proves the trailing 0x55 really was the bucket
// frame's terminator: Portisch builds pulse entries from alternating
// signal edges, so the two level bits inside one pulse byte are always
// opposite — 0xAA (two high-level nibbles) cannot occur in pulse data.
// Finalize before this byte starts the new frame, so back-to-back
// deliveries are split even when loop() never observed a quiet gap
// between them.
this->finish_bucket_frame_();
}
size_t at = this->rx_buffer_.size();
this->rx_buffer_.push_back(byte);
const uint8_t *raw = &this->rx_buffer_[0];
@@ -84,26 +94,21 @@ bool RFBridgeComponent::parse_bridge_byte_(uint8_t byte) {
break;
}
case RF_CODE_RFIN_BUCKET: {
if (byte != RF_CODE_STOP) {
return true;
if (at == 2) {
// The count byte: Portisch sends at most 7 buckets + sync, so 0 or
// >8 cannot be a genuine capture — reject before it can occupy the
// buffer for a full frame timeout.
return byte != 0 && byte <= B1_MAX_BUCKET_COUNT;
}
uint8_t buckets = raw[2] << 1;
std::string str;
char next_byte[3]; // 2 hex chars + null
for (uint32_t i = 0; i <= at; i++) {
buf_append_printf(next_byte, sizeof(next_byte), 0, "%02X", raw[i]);
str += next_byte;
if ((i > 3) && buckets) {
buckets--;
}
if ((i < 3) || (buckets % 2) || (i == at - 1)) {
str += " ";
}
}
ESP_LOGI(TAG, "Received RFBridge Bucket: %s", str.c_str());
break;
// 0x55 is legal DATA inside a B1 frame: bucket durations are sent
// with only their HIGH byte masked to 7 bits, so a duration such as
// 0x0155 puts a raw 0x55 low byte inside the table — the first 0x55
// must therefore not end the capture. The header declares the table
// length (raw[2] pairs), so a 0x55 there is always data; one at or
// past the first pulse index is a terminator CANDIDATE, confirmed
// once the UART goes quiet (finish_bucket_frame_ in loop()).
this->bucket_frame_candidate_ = byte == RF_CODE_STOP && at >= 3 + static_cast<size_t>(raw[2]) * 2;
return true;
}
default:
ESP_LOGW(TAG, "Unknown action: 0x%02X", action);
@@ -119,6 +124,47 @@ bool RFBridgeComponent::parse_bridge_byte_(uint8_t byte) {
return false;
}
void RFBridgeComponent::finish_bucket_frame_() {
if (this->rx_buffer_.size() < 4) {
// The candidate flag requires a header + non-empty bucket table, so
// this cannot happen while flag and buffer stay consistent; guard the
// raw[2] / size-1 reads against any future divergence anyway.
this->rx_buffer_.clear();
this->bucket_frame_candidate_ = false;
return;
}
const uint8_t *raw = this->rx_buffer_.data();
const size_t at = this->rx_buffer_.size() - 1;
uint8_t buckets = raw[2] << 1;
std::string str;
char next_byte[3]; // 2 hex chars + null
for (uint32_t i = 0; i <= at; i++) {
buf_append_printf(next_byte, sizeof(next_byte), 0, "%02X", raw[i]);
str += next_byte;
if ((i > 3) && buckets) {
buckets--;
}
if ((i < 3) || (buckets % 2) || (i == at - 1)) {
str += " ";
}
}
ESP_LOGI(TAG, "Received RFBridge Bucket: %s", str.c_str());
// Deliberately NOT ACKed: Portisch's B1 command handler leaves its
// last_sniffing_command at the previous mode (RF_CODE_RFIN), and its
// host-ACK handler re-arms sniffing from that stale value — so ACKing a
// bucket delivery silently reverts the radio to standard sniffing and
// ends bucket capture. Its delivery path is fire-and-forget and never
// waits for a host ACK. Stock Itead firmware never sends B1 frames, so
// suppressing this ACK cannot change stock-firmware behavior.
// https://github.com/esphome/esphome/issues/17682
this->rx_buffer_.clear();
this->bucket_frame_candidate_ = false;
}
void RFBridgeComponent::write_byte_str_(const std::string &codes) {
uint8_t code;
int size = codes.length();
@@ -130,12 +176,31 @@ void RFBridgeComponent::write_byte_str_(const std::string &codes) {
void RFBridgeComponent::loop() {
const uint32_t now = App.get_loop_component_start_time();
if (now - this->last_bridge_byte_ > 50) {
size_t avail = this->available();
if (avail == 0 && this->bucket_frame_candidate_ && now - this->last_bridge_byte_ > BUCKET_CANDIDATE_QUIET_MS) {
// The trailing 0x55 was followed by UART quiet, so it really was the
// frame terminator and not an interior data byte.
this->finish_bucket_frame_();
this->last_bridge_byte_ = now;
}
const bool receiving_bucket = this->rx_buffer_.size() >= 2 && this->rx_buffer_[1] == RF_CODE_RFIN_BUCKET;
if (receiving_bucket) {
// Never declare an in-progress bucket frame dead while its continuation
// bytes are already queued: a stalled loop() otherwise discards a live
// frame that the UART buffer proves is still arriving.
if (avail == 0 && now - this->last_bridge_byte_ > BUCKET_FRAME_TIMEOUT_MS) {
ESP_LOGD(TAG, "Discarding incomplete RFBridge Bucket frame (%u bytes)",
static_cast<unsigned>(this->rx_buffer_.size()));
this->rx_buffer_.clear();
this->bucket_frame_candidate_ = false;
this->last_bridge_byte_ = now;
}
} else if (now - this->last_bridge_byte_ > 50) {
this->rx_buffer_.clear();
this->bucket_frame_candidate_ = false;
this->last_bridge_byte_ = now;
}
size_t avail = this->available();
while (avail > 0) {
uint8_t buf[64];
size_t to_read = std::min(avail, sizeof(buf));
@@ -146,12 +211,14 @@ void RFBridgeComponent::loop() {
for (size_t i = 0; i < to_read; i++) {
if (this->rx_buffer_.size() > MAX_RX_BUFFER_SIZE) {
this->rx_buffer_.clear();
this->bucket_frame_candidate_ = false;
}
if (this->parse_bridge_byte_(buf[i])) {
ESP_LOGVV(TAG, "Parsed: 0x%02X", buf[i]);
this->last_bridge_byte_ = now;
} else {
this->rx_buffer_.clear();
this->bucket_frame_candidate_ = false;
}
}
}
+13
View File
@@ -30,6 +30,17 @@ static const uint8_t RF_CODE_BEEP = 0xC0;
static const uint8_t RF_CODE_STOP = 0x55;
static const uint8_t RF_DEBOUNCE = 200;
static const size_t MAX_RX_BUFFER_SIZE = 512;
// ~10 byte times at 19200 baud: long enough to prove the UART went quiet
// after a possible bucket-frame terminator, short enough to finish well
// before the next radio capture can be delivered.
static const uint32_t BUCKET_CANDIDATE_QUIET_MS = 5;
// Portisch drains a B1 frame's header, bucket table, and pulse data as
// separate UART writes, so an in-progress bucket frame tolerates a longer
// inter-region gap than the generic 50 ms inter-byte timeout.
static const uint32_t BUCKET_FRAME_TIMEOUT_MS = 250;
// Portisch's uart_put_RF_buckets sends at most 7 buckets plus the sync
// bucket, so a B1 count byte above 8 (or 0) is malformed for any protocol.
static const uint8_t B1_MAX_BUCKET_COUNT = 8;
struct RFBridgeData {
uint16_t sync;
@@ -67,10 +78,12 @@ class RFBridgeComponent final : public uart::UARTDevice, public Component {
void ack_();
void decode_();
bool parse_bridge_byte_(uint8_t byte);
void finish_bucket_frame_();
void write_byte_str_(const std::string &codes);
std::vector<uint8_t> rx_buffer_;
uint32_t last_bridge_byte_{0};
bool bucket_frame_candidate_{false};
CallbackManager<void(RFBridgeData)> data_callback_;
CallbackManager<void(RFBridgeAdvancedData)> advanced_data_callback_;
+14 -3
View File
@@ -1,10 +1,13 @@
#include "tuya.h"
#include "esphome/components/network/util.h"
#include "esphome/core/gpio.h"
#include "esphome/core/helpers.h"
#include "esphome/core/log.h"
#include "esphome/core/util.h"
#ifdef USE_NETWORK
#include "esphome/components/network/util.h"
#endif
#ifdef USE_WIFI
#include "esphome/components/wifi/wifi_component.h"
#endif
@@ -22,6 +25,14 @@ static const int MAX_RETRIES = 5;
// Max bytes to log for datapoint values (larger values are truncated)
static constexpr size_t MAX_DATAPOINT_LOG_BYTES = 16;
static bool network_is_connected() {
#ifdef USE_NETWORK
return network::is_connected();
#else
return false;
#endif
}
void Tuya::setup() {
this->set_interval("heartbeat", 15000, [this] { this->send_empty_command_(TuyaCommandType::HEARTBEAT); });
if (this->status_pin_ != nullptr) {
@@ -554,14 +565,14 @@ void Tuya::send_empty_command_(TuyaCommandType command) {
}
void Tuya::set_status_pin_() {
bool is_network_ready = network::is_connected() && remote_is_connected();
bool is_network_ready = network_is_connected() && remote_is_connected();
this->status_pin_->digital_write(is_network_ready);
}
uint8_t Tuya::get_wifi_status_code_() {
uint8_t status = 0x02;
if (network::is_connected()) {
if (network_is_connected()) {
status = 0x03;
// Protocol version 3 also supports specifying when connected to "the cloud"
+4 -12
View File
@@ -1,5 +1,4 @@
from collections.abc import Callable
from typing import Any, NoReturn
from typing import Any
from esphome import automation
from esphome.automation import Trigger
@@ -48,17 +47,10 @@ UDP_SCHEMA = cv.Schema(
)
def is_relocated(option: str) -> Callable[[Any], NoReturn]:
def validator(value: Any) -> NoReturn:
raise cv.Invalid(
f"The '{option}' option should now be configured in the 'packet_transport' component"
)
return validator
RELOCATED = {
cv.Optional(x): is_relocated(x)
cv.Optional(x): cv.invalid(
f"The '{x}' option should now be configured in the 'packet_transport' component"
)
for x in (
CONF_PROVIDERS,
CONF_ENCRYPTION,
+4 -16
View File
@@ -22,10 +22,6 @@ namespace esphome {
* pointer. When it is default constructed, it has empty string. You can freely copy or move around this struct, but
* never free its pointer. str() function can be used to export the content as std::string. StringRef is adopted from
* <https://github.com/nghttp2/nghttp2/blob/29cbf8b83ff78faf405d1086b16adc09a8772eca/src/template.h#L376>
*
* A StringRef may carry a null pointer while its length is zero (the generated api messages start their encode only
* string fields that way). Every member treats that as the empty string; only c_str() hands the null pointer on, so
* callers that print or copy through c_str() must check empty() first.
*/
class StringRef {
public:
@@ -82,7 +78,7 @@ class StringRef {
/// True if the view begins with the given prefix (std::string::starts_with-like)
bool starts_with(const StringRef &prefix) const {
return len_ >= prefix.len_ && (prefix.len_ == 0 || std::memcmp(base_, prefix.base_, prefix.len_) == 0);
return len_ >= prefix.len_ && std::memcmp(base_, prefix.base_, prefix.len_) == 0;
}
bool starts_with(const char *prefix) const { return this->starts_with(StringRef(prefix)); }
bool starts_with(const std::string &prefix) const { return this->starts_with(StringRef(prefix)); }
@@ -96,15 +92,14 @@ class StringRef {
return actual;
}
std::string str() const { return std::string(base_, len_); } // fine for {nullptr, 0}: nothing is read
std::string str() const { return std::string(base_, len_); }
const uint8_t *byte() const { return reinterpret_cast<const uint8_t *>(base_); }
operator std::string() const { return str(); }
/// Compare (compatible with std::string::compare)
int compare(const StringRef &other) const {
size_type common = std::min(len_, other.len_);
int result = common == 0 ? 0 : std::memcmp(base_, other.base_, common);
int result = std::memcmp(base_, other.base_, std::min(len_, other.len_));
if (result != 0)
return result;
if (len_ < other.len_)
@@ -263,14 +258,7 @@ inline double stod(const StringRef &str, size_t *pos = nullptr) {
#ifdef USE_JSON
// NOLINTNEXTLINE(readability-identifier-naming)
inline void convertToJson(const StringRef &src, JsonVariant dst) {
// Bounded by the view length; a null, empty view becomes "" rather than JSON null
if (src.empty()) {
dst.set("");
return;
}
dst.set(JsonString(src.c_str(), src.size()));
}
inline void convertToJson(const StringRef &src, JsonVariant dst) { dst.set(src.c_str()); }
#endif // USE_JSON
} // namespace esphome
+1 -4
View File
@@ -69,10 +69,7 @@ def _make_create_connection() -> Callable[..., socket.socket]:
from aiohappyeyeballs import start_connection
from urllib3.exceptions import LocationParseError
from urllib3.util.connection import ( # noqa: PLC2701
_set_socket_options,
allowed_gai_family,
)
from urllib3.util.connection import _set_socket_options, allowed_gai_family # noqa: PLC2701
from urllib3.util.timeout import _DEFAULT_TIMEOUT # noqa: PLC2701
from esphome import async_thread
+5 -1
View File
@@ -616,11 +616,15 @@ def _make_registry_client() -> Any:
elsewhere, not by the PlatformIO registry.
"""
from platformio.package.manager._registry import PackageManagerRegistryMixin
from platformio.registry.client import RegistryClient
class _Registry(PackageManagerRegistryMixin):
def __init__(self) -> None:
self._registry_client = None
self.pkg_type = "library"
self._registry_client = RegistryClient()
# The probe sleeps ~500 ms per lookup (see runner.patch_registry_private_packages);
# instance-level so the ESPHome process never patches PlatformIO's class
self._registry_client.allowed_private_packages = lambda: False
@staticmethod
def is_system_compatible(value: Any, custom_system: Any = None) -> bool:
+2
View File
@@ -951,8 +951,10 @@ def main(argv: list[str]) -> int:
"""Subprocess entry point: ``prefetch <build_dir> <env_name>``."""
from esphome.core import CORE
from esphome.log import setup_log
from esphome.platformio.runner import patch_registry_private_packages
signal.signal(signal.SIGTERM, _sigterm)
patch_registry_private_packages()
raw_level = os.environ.get("ESPHOME_PREFETCH_LOG_LEVEL")
try:
level = int(raw_level) if raw_level is not None else logging.INFO
+13 -1
View File
@@ -2,7 +2,8 @@
Invoked via ``python -m esphome.platformio.runner`` instead of
``python -m platformio`` so that the patches (incremental rebuild
preservation, download retries) apply inside the subprocess. Running
preservation, download retries, skipping the private-package probe) apply
inside the subprocess. Running
PlatformIO in a subprocess keeps its ``sys.path`` mutations and other
global state from leaking into the ESPHome process.
"""
@@ -105,6 +106,16 @@ def patch_file_downloader() -> None:
FileDownloader.__init__ = patched_init
def patch_registry_private_packages() -> None:
"""Skip PlatformIO's private-package probe; it sleeps ~500 ms per lookup.
ESPHome never uses private packages, so the answer is always False.
"""
from platformio.registry.client import RegistryClient
RegistryClient.allowed_private_packages = staticmethod(lambda: False) # type: ignore[method-assign]
_IGNORE_LIB_WARNINGS = "(?:Hash|Update)"
# Regex patterns matched against each line of PlatformIO output. Lines that
# match are dropped by RedirectText before they reach the parent process.
@@ -152,6 +163,7 @@ FILTER_PLATFORMIO_LINES = [
def main() -> int:
patch_structhash()
patch_file_downloader()
patch_registry_private_packages()
# Wrap stdout/stderr with RedirectText before PlatformIO runs:
#
+2 -2
View File
@@ -1,4 +1,4 @@
# Useful stuff when working in a development environment
clang-format==13.0.1 # also change in .pre-commit-config.yaml and Dockerfile when updating
clang-format==13.0.1 # .pre-commit-config.yaml rev synced by script/sync_dependency_versions.py
clang-tidy==22.1.8
yamllint==1.38.0 # also change in .pre-commit-config.yaml when updating
yamllint==1.38.0 # .pre-commit-config.yaml rev synced by script/sync_dependency_versions.py
+5 -4
View File
@@ -1,8 +1,9 @@
pylint==4.0.8
flake8==7.3.0 # also change in .pre-commit-config.yaml when updating
ruff==0.16.5 # also change in .pre-commit-config.yaml when updating
pyupgrade==3.21.2 # also change in .pre-commit-config.yaml when updating
prek==0.5.1 # also change in .github/workflows/ci.yml when updating
flake8==7.3.0 # .pre-commit-config.yaml rev synced by script/sync_dependency_versions.py
ruff==0.16.6 # .pre-commit-config.yaml rev synced by script/sync_dependency_versions.py
pyupgrade==3.21.2 # .pre-commit-config.yaml rev synced by script/sync_dependency_versions.py
prek==0.5.2 # .github/workflows/ci.yml reads this pin
yamlrocks==0.6.1 # used by script/sync_dependency_versions.py
# Unit tests
pytest==9.1.1
+314 -239
View File
@@ -28,11 +28,6 @@ class WireType(IntEnum):
END_GROUP = 4 # groups (deprecated)
FIXED32 = 5 # fixed32, sfixed32, float
@property
def cpp_name(self) -> str:
"""The matching constant in proto.h."""
return f"WIRE_TYPE_{self.name}"
# Generate with
# protoc --python_out=script/api_protobuf -I esphome/components/api/ api_options.proto
@@ -131,10 +126,9 @@ def camel_to_snake(name: str) -> str:
return re.sub("([a-z0-9])([A-Z])", r"\1_\2", s1).lower()
def _encode_call(func: str, *args: str, force: bool = False) -> str:
"""Emit one ProtoEncode call; every helper takes the cursor and returns it advanced."""
suffix = "_force" if force else ""
return f"pos = ProtoEncode::{func}{suffix}({', '.join(('pos', *args))});"
def force_str(force: bool) -> str:
"""Convert a boolean force value to string format for C++ code."""
return str(force).lower()
class TypeInfo(ABC):
@@ -229,39 +223,55 @@ class TypeInfo(ABC):
def class_member(self) -> str:
return f"{self.cpp_type} {self.field_name}{{{self.default_value}}};"
def decode_case(self, body: str) -> str:
"""Emit one decode_field() case, keyed on the field's wire tag."""
return f"case proto_tag({self.number}, {self.wire_type.cpp_name}):\n" + indent(
f"{body}\nbreak;"
)
@property
def decode_varint_content(self) -> str:
content = self.decode_varint
if content is None:
return None
return f"case {self.number}: this->{self.field_name} = {content}; break;"
# Expression that reads this field from `value`; None when the type is never decoded.
decode_expr: str | None = None
def _decode_store(self, expr: str) -> str:
return f"this->{self.field_name} = {expr};"
decode_varint = None
@property
def decode_content(self) -> str | None:
"""The decode_field() case for this field, or None when it is never decoded."""
expr = self.decode_expr
return None if expr is None else self.decode_case(self._decode_store(expr))
def decode_length_content(self) -> str:
content = self.decode_length
if content is None:
return None
return f"case {self.number}: this->{self.field_name} = {content}; break;"
decode_length = None
@property
def decode_32bit_content(self) -> str:
content = self.decode_32bit
if content is None:
return None
return f"case {self.number}: this->{self.field_name} = {content}; break;"
decode_32bit = None
@property
def decode_64bit_content(self) -> str:
content = self.decode_64bit
if content is None:
return None
return f"case {self.number}: this->{self.field_name} = {content}; break;"
decode_64bit = None
# Mapping from encode_func to raw encode expression template.
# When a forced field has a single-byte tag, the code generator emits
# write_raw_byte(tag) + raw encode instead of the full encode_* method,
# eliminating the zero-check branch and encode_field_raw indirection.
# {value} is replaced with the actual field expression.
RAW_ENCODE_MAP: dict[str, tuple[str, str]] = {
"encode_uint32": ("encode_varint_raw", "{value}"),
"encode_uint64": ("encode_varint_raw_64", "{value}"),
"encode_sint32": ("encode_varint_raw_short", "encode_zigzag32({value})"),
"encode_sint64": ("encode_varint_raw_64", "encode_zigzag64({value})"),
"encode_int64": ("encode_varint_raw_64", "static_cast<uint64_t>({value})"),
"encode_bool": ("write_raw_byte", "{value} ? 0x01 : 0x00"),
RAW_ENCODE_MAP: dict[str, str] = {
"encode_uint32": "ProtoEncode::encode_varint_raw(pos, {value});",
"encode_uint64": "ProtoEncode::encode_varint_raw_64(pos, {value});",
"encode_sint32": "ProtoEncode::encode_varint_raw_short(pos, encode_zigzag32({value}));",
"encode_sint64": "ProtoEncode::encode_varint_raw_64(pos, encode_zigzag64({value}));",
"encode_int64": "ProtoEncode::encode_varint_raw_64(pos, static_cast<uint64_t>({value}));",
"encode_bool": "ProtoEncode::write_raw_byte(pos, {value} ? 0x01 : 0x00);",
}
# Fixed32 value expression for the shared tag+fixed32 writer; None for other wire types
fixed32_value_template: str | None = None
def _encode_with_precomputed_tag(self, value_expr: str) -> str | None:
"""Try to emit a precomputed-tag encode for a field.
@@ -278,17 +288,12 @@ class TypeInfo(ABC):
return None
max_val = self.max_value
# Only use RAW_ENCODE_MAP for forced fields or fields with max_value
raw = None
raw_expr = None
if self.force or max_val is not None:
raw = self.RAW_ENCODE_MAP.get(self.encode_func)
if raw is None:
raw_expr = self.RAW_ENCODE_MAP.get(self.encode_func)
if raw_expr is None:
return None
func, arg = raw
body = (
_encode_call("write_raw_byte", str(tag))
+ "\n"
+ _encode_call(func, arg.format(value=value_expr))
)
body = f"ProtoEncode::write_raw_byte(pos, {tag});\n{raw_expr.format(value=value_expr)}"
if self.force:
return body
# Non-forced with max_value: inline zero-check + raw encode
@@ -309,44 +314,23 @@ class TypeInfo(ABC):
return None
# When max_len < 128, length varint is always 1 byte
len_encode = (
_encode_call("write_raw_byte", f"static_cast<uint8_t>({len_expr})")
f"ProtoEncode::write_raw_byte(pos, static_cast<uint8_t>({len_expr}));"
if max_len is not None and max_len < 128
else _encode_call("encode_varint_raw", len_expr)
else f"ProtoEncode::encode_varint_raw(pos, {len_expr});"
)
return "\n".join(
(
_encode_call("write_raw_byte", str(tag)),
len_encode,
_encode_call("encode_raw", data_expr, len_expr),
)
)
def _encode_fixed32_with_precomputed_tag(self, value: str) -> str | None:
"""Single-byte tag fixed32 write, or None for other types and multi-byte tags."""
tag = self.calculate_tag()
if self.fixed32_value_template is None or tag >= 128:
return None
value_expr = self.fixed32_value_template.format(value=value)
if self.force:
return _encode_call("write_tag_and_fixed32", str(tag), value_expr)
return (
f"if (uint32_t raw = {value_expr}; raw != 0) [[likely]] {{\n"
f" {_encode_call('write_tag_and_fixed32', str(tag), 'raw')}\n"
"}"
f"ProtoEncode::write_raw_byte(pos, {tag});\n"
f"{len_encode}\n"
f"ProtoEncode::encode_raw(pos, {data_expr}, {len_expr});"
)
@property
def encode_content(self) -> str:
value = f"this->{self.field_name}"
if result := self._encode_with_precomputed_tag(value):
if result := self._encode_with_precomputed_tag(f"this->{self.field_name}"):
return result
if result := self._encode_fixed32_with_precomputed_tag(value):
return result
return _encode_call(self.encode_func, str(self.number), value, force=self.force)
def encode_element(self, number: int, element: str) -> str:
"""Encode one element of a repeated field; elements are always written."""
return _encode_call(self.encode_func, str(number), element, force=True)
if self.force:
return f"ProtoEncode::{self.encode_func}(pos, {self.number}, this->{self.field_name}, true);"
return f"ProtoEncode::{self.encode_func}(pos, {self.number}, this->{self.field_name});"
encode_func = None
@@ -566,17 +550,17 @@ def create_field_type_info(
# For messages that decode (SOURCE_CLIENT or SOURCE_BOTH), use pointer
# for zero-copy access to the receive buffer
if needs_decode:
return PointerToBytesBufferType(field, needs_decode)
return PointerToBytesBufferType(field, None)
# For SOURCE_SERVER (encode only), explicit annotation is still needed
if get_field_opt(field, pb.pointer_to_buffer, False):
return PointerToBytesBufferType(field, needs_decode)
return PointerToBytesBufferType(field, None)
return BytesType(field, needs_decode, needs_encode)
# Special handling for string fields - use StringRef for zero-copy
if field.type == 9:
return PointerToStringBufferType(field, needs_decode)
return PointerToStringBufferType(field, None)
validate_field_type(field.type, field.name)
if field.type == 11:
@@ -621,6 +605,7 @@ class DoubleType(FixedSizeTypeMixin, TypeInfo):
# Unsupported but defined for completeness
cpp_type = "double"
default_value = "0.0"
decode_64bit = "value.as_double()"
encode_func = "encode_double"
wire_type = WireType.FIXED64 # Uses wire type 1 according to protobuf spec
@@ -646,12 +631,10 @@ class DoubleType(FixedSizeTypeMixin, TypeInfo):
class FloatType(FixedSizeTypeMixin, TypeInfo):
cpp_type = "float"
default_value = "0.0f"
decode_expr = "value.as_float()"
decode_32bit = "value.as_float()"
encode_func = "encode_float"
wire_type = WireType.FIXED32 # Uses wire type 5
fixed32_value_template = "float_to_raw({value})"
def dump(self, name: str) -> str:
o = f'snprintf(buffer, sizeof(buffer), "%g", {name});\n'
o += "out.append(buffer);"
@@ -675,7 +658,7 @@ class Int64Type(VarintTypeMixin, TypeInfo):
cpp_type = "int64_t"
_varint_max_bits = 64
default_value = "0"
decode_expr = "static_cast<int64_t>(value.as_varint())"
decode_varint = "static_cast<int64_t>(value)"
encode_func = "encode_int64"
wire_type = WireType.VARINT # Uses wire type 0
@@ -696,7 +679,7 @@ class UInt64Type(VarintTypeMixin, TypeInfo):
cpp_type = "uint64_t"
_varint_max_bits = 64
default_value = "0"
decode_expr = "value.as_varint()"
decode_varint = "value"
encode_func = "encode_uint64"
wire_type = WireType.VARINT # Uses wire type 0
@@ -714,11 +697,11 @@ class UInt64Type(VarintTypeMixin, TypeInfo):
return self._get_simple_size_calculation(name, force, "uint64")
@property
def RAW_ENCODE_MAP(self) -> dict[str, tuple[str, str]]: # noqa: N802
def RAW_ENCODE_MAP(self) -> dict[str, str]: # noqa: N802
if self.mac_address:
return {
**TypeInfo.RAW_ENCODE_MAP,
"encode_uint64": ("encode_varint_raw_48bit", "{value}"),
"encode_uint64": "ProtoEncode::encode_varint_raw_48bit(pos, {value});",
}
return TypeInfo.RAW_ENCODE_MAP
@@ -731,7 +714,7 @@ class Int32Type(VarintTypeMixin, TypeInfo):
cpp_type = "int32_t"
_varint_max_bits = 64 # int32 is sign-extended to 64 bits in protobuf
default_value = "0"
decode_expr = "static_cast<int32_t>(value.as_varint())"
decode_varint = "static_cast<int32_t>(value)"
encode_func = "encode_int32"
wire_type = WireType.VARINT # Uses wire type 0
@@ -751,6 +734,7 @@ class Int32Type(VarintTypeMixin, TypeInfo):
class Fixed64Type(FixedSizeTypeMixin, TypeInfo):
cpp_type = "uint64_t"
default_value = "0"
decode_64bit = "value.as_fixed64()"
encode_func = "encode_fixed64"
wire_type = WireType.FIXED64 # Uses wire type 1
@@ -776,7 +760,7 @@ class Fixed64Type(FixedSizeTypeMixin, TypeInfo):
class Fixed32Type(FixedSizeTypeMixin, TypeInfo):
cpp_type = "uint32_t"
default_value = "0"
decode_expr = "value.as_fixed32()"
decode_32bit = "value.as_fixed32()"
encode_func = "encode_fixed32"
wire_type = WireType.FIXED32 # Uses wire type 5
@@ -785,7 +769,15 @@ class Fixed32Type(FixedSizeTypeMixin, TypeInfo):
o += "out.append(buffer);"
return o
fixed32_value_template = "{value}"
@property
def encode_content(self) -> str:
tag = self.calculate_tag()
if self.force and tag < 128:
# Emit combined tag+value write: precomputed tag + direct memcpy
return f"ProtoEncode::write_tag_and_fixed32(pos, {tag}, this->{self.field_name});"
if self.force:
return f"ProtoEncode::{self.encode_func}(pos, {self.number}, this->{self.field_name}, true);"
return f"ProtoEncode::{self.encode_func}(pos, {self.number}, this->{self.field_name});"
def get_size_calculation(self, name: str, force: bool = False) -> str:
field_id_size = self.calculate_field_id_size()
@@ -805,7 +797,7 @@ class BoolType(VarintTypeMixin, TypeInfo):
_varint_max_bits = 1
cpp_type = "bool"
default_value = "false"
decode_expr = "value.as_bool()"
decode_varint = "value != 0"
encode_func = "encode_bool"
wire_type = WireType.VARINT # Uses wire type 0
@@ -825,7 +817,7 @@ class StringType(TypeInfo):
default_value = ""
reference_type = "std::string &"
const_reference_type = "const std::string &"
decode_expr = "value.as_string()"
decode_length = "value.as_string()"
encode_func = "encode_string"
wire_type = WireType.LENGTH_DELIMITED # Uses wire type 2
@@ -859,12 +851,9 @@ class StringType(TypeInfo):
f"this->{self.field_name}_ref_.size()",
):
return result
return _encode_call(
"encode_string",
str(self.number),
f"this->{self.field_name}_ref_",
force=self.force,
)
if self.force:
return f"ProtoEncode::encode_string(pos, {self.number}, this->{self.field_name}_ref_, true);"
return f"ProtoEncode::encode_string(pos, {self.number}, this->{self.field_name}_ref_);"
def dump(self, name):
# If name is 'it', this is a repeated field element - always use string
@@ -940,9 +929,6 @@ class MessageType(TypeInfo):
def can_use_dump_field(cls) -> bool:
return False
def encode_element(self, number: int, element: str) -> str:
return _encode_call("encode_sub_message", "buffer", str(number), element)
@property
def cpp_type(self) -> str:
return self._field.type_name[1:]
@@ -965,9 +951,15 @@ class MessageType(TypeInfo):
@property
def encode_content(self) -> str:
# Sub-message encoding needs buffer for backpatch/sync
return _encode_call(
self.encode_func, "buffer", str(self.number), f"this->{self.field_name}"
)
return f"ProtoEncode::{self.encode_func}(pos, buffer, {self.number}, this->{self.field_name});"
@property
def decode_length(self) -> str:
# Override to return None for message types because we can't use template-based
# decoding when the specific message type isn't known at compile time.
# Instead, we use the non-template decode_to_message() method which allows
# runtime polymorphism through virtual function calls.
return None
@property
def public_content(self) -> list[str]:
@@ -984,14 +976,19 @@ class MessageType(TypeInfo):
)
@property
def decode_content(self) -> str:
body = f"value.decode_to_message(this->{self.field_name});"
def decode_length_content(self) -> str:
# Custom decode that doesn't use templates
if self._track_presence:
# decode_to_message() cannot report failure, so setting the flag
# afterwards only documents intent; a status-returning decode could
# gate it for real without touching callers.
body += f"\nthis->has_{self.name} = true;"
return self.decode_case(body)
return (
f"case {self.number}:\n"
f" value.decode_to_message(this->{self.field_name});\n"
f" this->has_{self.name} = true;\n"
f" break;"
)
return f"case {self.number}: value.decode_to_message(this->{self.field_name}); break;"
def dump(self, name: str) -> str:
return f"{name}.dump_to(out);"
@@ -1030,7 +1027,7 @@ class BytesType(TypeInfo):
reference_type = "std::string &"
const_reference_type = "const std::string &"
encode_func = "encode_bytes"
decode_expr = "value.as_string()"
decode_length = "value.as_string()"
wire_type = WireType.LENGTH_DELIMITED # Uses wire type 2
@property
@@ -1061,13 +1058,9 @@ class BytesType(TypeInfo):
f"this->{self.field_name}_ptr_", f"this->{self.field_name}_len_"
):
return result
return _encode_call(
"encode_bytes",
str(self.number),
f"this->{self.field_name}_ptr_",
f"this->{self.field_name}_len_",
force=self.force,
)
if self.force:
return f"ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}_ptr_, this->{self.field_name}_len_, true);"
return f"ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}_ptr_, this->{self.field_name}_len_);"
def dump(self, name: str) -> str:
ptr_dump = f"format_hex_pretty(this->{self.field_name}_ptr_, this->{self.field_name}_len_)"
@@ -1134,12 +1127,16 @@ class PointerToBufferTypeBase(TypeInfo):
def can_use_dump_field(cls) -> bool:
return False
# Only here to make needs_decode required: the null string default keys off it, so a call
# site must not fall back on the base class default
def __init__(
self, field: descriptor.FieldDescriptorProto, needs_decode: bool
self, field: descriptor.FieldDescriptorProto, size: int | None = None
) -> None:
super().__init__(field, needs_decode)
super().__init__(field)
self.array_size = 0
@property
def decode_length(self) -> str | None:
# This is handled in decode_length_content
return None
@property
def wire_type(self) -> WireType:
@@ -1173,20 +1170,17 @@ class PointerToBytesBufferType(PointerToBufferTypeBase):
f"this->{self.field_name}", f"this->{self.field_name}_len"
):
return result
return _encode_call(
"encode_bytes",
str(self.number),
f"this->{self.field_name}",
f"this->{self.field_name}_len",
force=self.force,
)
if self.force:
return f"ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}, this->{self.field_name}_len, true);"
return f"ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}, this->{self.field_name}_len);"
@property
def decode_content(self) -> str:
return self.decode_case(
f"this->{self.field_name} = value.data();\n"
f"this->{self.field_name}_len = value.size();",
)
def decode_length_content(self) -> str | None:
return f"""case {self.number}: {{
this->{self.field_name} = value.data();
this->{self.field_name}_len = value.size();
break;
}}"""
def dump(self, name: str) -> str:
return (
@@ -1220,53 +1214,34 @@ class PointerToStringBufferType(PointerToBufferTypeBase):
def can_use_dump_field(cls) -> bool:
return True
@property
def _starts_null(self) -> bool:
"""A field that is only encoded, and skipped when empty, never has its pointer read
before it is set, so it can default to a null StringRef and the message constructs as
one zero fill. Any encode path that copies unconditionally must check this."""
return not self._needs_decode and not self.force
@property
def public_content(self) -> list[str]:
if self._starts_null:
return [
f"StringRef {self.field_name}{{nullptr, 0}}; // null until set, encode only"
]
return [f"StringRef {self.field_name}{{}};"]
@property
def encode_content(self) -> str:
max_len = self.max_data_length
if max_len is not None and max_len < 128 and self.force:
assert not self._starts_null, (
"unconditional copy of a field that may start null"
)
tag = self.calculate_tag()
if tag < 128:
return _encode_call(
"encode_short_string_force", str(tag), f"this->{self.field_name}"
)
return f"ProtoEncode::encode_short_string_force(pos, {tag}, this->{self.field_name});"
if result := self._encode_bytes_with_precomputed_tag(
f"this->{self.field_name}.c_str()",
f"this->{self.field_name}.size()",
):
assert not self._starts_null, (
"unconditional copy of a field that may start null"
)
return result
return _encode_call(
"encode_string",
str(self.number),
f"this->{self.field_name}",
force=self.force,
if self.force:
return f"ProtoEncode::encode_string(pos, {self.number}, this->{self.field_name}, true);"
return (
f"ProtoEncode::encode_string(pos, {self.number}, this->{self.field_name});"
)
@property
def decode_content(self) -> str:
return self.decode_case(
f"this->{self.field_name} = StringRef(value.data(), value.size());",
)
def decode_length_content(self) -> str | None:
return f"""case {self.number}: {{
this->{self.field_name} = StringRef(reinterpret_cast<const char *>(value.data()), value.size());
break;
}}"""
def dump(self, name: str) -> str:
# Not used since we use dump_field, but required by abstract base class
@@ -1335,13 +1310,14 @@ class PackedBufferTypeInfo(TypeInfo):
]
@property
def decode_content(self) -> str:
def decode_length_content(self) -> str:
"""Store pointer to buffer and calculate count of packed varints."""
return self.decode_case(
f"this->{self.field_name}_data_ = value.data();\n"
f"this->{self.field_name}_length_ = value.size();\n"
f"this->{self.field_name}_count_ = count_packed_varints(value.data(), value.size());",
)
return f"""case {self.number}: {{
this->{self.field_name}_data_ = value.data();
this->{self.field_name}_length_ = value.size();
this->{self.field_name}_count_ = count_packed_varints(value.data(), value.size());
break;
}}"""
@property
def encode_content(self) -> str:
@@ -1426,11 +1402,17 @@ class FixedArrayBytesType(TypeInfo):
]
@property
def decode_content(self) -> str:
return self.decode_case(
f"this->{self.field_name}_len = std::min<size_t>(value.size(), {self.array_size});\n"
f"memcpy(this->{self.field_name}, value.data(), this->{self.field_name}_len);",
)
def decode_length_content(self) -> str:
o = f"case {self.number}: {{\n"
o += " const std::string &data_str = value.as_string();\n"
o += f" this->{self.field_name}_len = data_str.size();\n"
o += f" if (this->{self.field_name}_len > {self.array_size}) {{\n"
o += f" this->{self.field_name}_len = {self.array_size};\n"
o += " }\n"
o += f" memcpy(this->{self.field_name}, data_str.data(), this->{self.field_name}_len);\n"
o += " break;\n"
o += "}"
return o
@property
def encode_content(self) -> str:
@@ -1439,13 +1421,9 @@ class FixedArrayBytesType(TypeInfo):
f"this->{self.field_name}", f"this->{self.field_name}_len", max_len=max_len
):
return result
return _encode_call(
"encode_bytes",
str(self.number),
f"this->{self.field_name}",
f"this->{self.field_name}_len",
force=self.force,
)
if self.force:
return f"ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}, this->{self.field_name}_len, true);"
return f"ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}, this->{self.field_name}_len);"
def dump(self, name: str) -> str:
return f"out.append(format_hex_pretty({name}, {name}_len));"
@@ -1493,7 +1471,7 @@ class UInt32Type(VarintTypeMixin, TypeInfo):
cpp_type = "uint32_t"
_varint_max_bits = 32
default_value = "0"
decode_expr = "value.as_varint()"
decode_varint = "value"
encode_func = "encode_uint32"
wire_type = WireType.VARINT # Uses wire type 0
@@ -1516,21 +1494,13 @@ class UInt32Type(VarintTypeMixin, TypeInfo):
class EnumType(VarintTypeMixin, TypeInfo):
_varint_max_bits = 32
def encode_element(self, number: int, element: str) -> str:
return _encode_call(
self.encode_func,
str(number),
f"static_cast<uint32_t>({element})",
force=True,
)
@property
def cpp_type(self) -> str:
return f"enums::{self._field.type_name[1:]}"
@property
def decode_expr(self) -> str:
return f"static_cast<{self.cpp_type}>(value.as_varint())"
def decode_varint(self) -> str:
return f"static_cast<{self.cpp_type}>(value)"
default_value = ""
wire_type = WireType.VARINT # Uses wire type 0
@@ -1550,9 +1520,9 @@ class EnumType(VarintTypeMixin, TypeInfo):
@property
def encode_content(self) -> str:
value_expr = f"static_cast<uint32_t>(this->{self.field_name})"
return _encode_call(
self.encode_func, str(self.number), value_expr, force=self.force
)
if self.force:
return f"ProtoEncode::{self.encode_func}(pos, {self.number}, {value_expr}, true);"
return f"ProtoEncode::{self.encode_func}(pos, {self.number}, {value_expr});"
def dump(self, name: str) -> str:
return f"out.append_p(proto_enum_to_string<{self.cpp_type}>({name}));"
@@ -1577,7 +1547,7 @@ class EnumType(VarintTypeMixin, TypeInfo):
class SFixed32Type(FixedSizeTypeMixin, TypeInfo):
cpp_type = "int32_t"
default_value = "0"
decode_expr = "value.as_sfixed32()"
decode_32bit = "value.as_sfixed32()"
encode_func = "encode_sfixed32"
wire_type = WireType.FIXED32 # Uses wire type 5
@@ -1603,6 +1573,7 @@ class SFixed32Type(FixedSizeTypeMixin, TypeInfo):
class SFixed64Type(FixedSizeTypeMixin, TypeInfo):
cpp_type = "int64_t"
default_value = "0"
decode_64bit = "value.as_sfixed64()"
encode_func = "encode_sfixed64"
wire_type = WireType.FIXED64 # Uses wire type 1
@@ -1629,7 +1600,7 @@ class SInt32Type(VarintTypeMixin, TypeInfo):
cpp_type = "int32_t"
_varint_max_bits = 32 # zigzag encoding keeps it 32-bit
default_value = "0"
decode_expr = "decode_zigzag32(static_cast<uint32_t>(value.as_varint()))"
decode_varint = "decode_zigzag32(static_cast<uint32_t>(value))"
encode_func = "encode_sint32"
wire_type = WireType.VARINT # Uses wire type 0
@@ -1650,7 +1621,7 @@ class SInt64Type(VarintTypeMixin, TypeInfo):
cpp_type = "int64_t"
_varint_max_bits = 64
default_value = "0"
decode_expr = "decode_zigzag64(value.as_varint())"
decode_varint = "decode_zigzag64(value)"
encode_func = "encode_sint64"
wire_type = WireType.VARINT # Uses wire type 0
@@ -1730,9 +1701,9 @@ def _generate_inline_encode_block(
lines = []
lines.append(f"auto &sub_msg = {element};")
lines.append(_encode_call("write_raw_byte", str(tag)))
lines.append(f"ProtoEncode::write_raw_byte(pos, {tag});")
lines.append("uint8_t *len_pos = pos;")
lines.append(_encode_call("reserve_byte"))
lines.append("ProtoEncode::reserve_byte(pos);")
# Generate inline field encoding for each sub-message field
for field in sub_desc.field:
@@ -1803,11 +1774,18 @@ class FixedArrayRepeatedType(TypeInfo):
def _encode_element(self, element: str) -> str:
"""Helper to generate encode statement for a single element."""
if isinstance(self._ti, MessageType) and _is_inline_encode(self._ti.cpp_type):
return _generate_inline_encode_block(
self.number, self._ti.cpp_type, element
)
return self._ti.encode_element(self.number, element)
if isinstance(self._ti, EnumType):
return f"ProtoEncode::{self._ti.encode_func}(pos, {self.number}, static_cast<uint32_t>({element}), true);"
# Repeated message elements use encode_sub_message (force=true is default)
if isinstance(self._ti, MessageType):
if _is_inline_encode(self._ti.cpp_type):
return _generate_inline_encode_block(
self.number, self._ti.cpp_type, element
)
return f"ProtoEncode::encode_sub_message(pos, buffer, {self.number}, {element});"
return (
f"ProtoEncode::{self._ti.encode_func}(pos, {self.number}, {element}, true);"
)
@property
def cpp_type(self) -> str:
@@ -2101,23 +2079,55 @@ class RepeatedTypeInfo(TypeInfo):
return self._ti.wire_type
@property
def decode_expr(self) -> str | None:
return self._ti.decode_expr
def _decode_store(self, expr: str) -> str:
return f"this->{self.field_name}.push_back({expr});"
@property
def decode_content(self) -> str | None:
def decode_varint_content(self) -> str:
# Pointer fields don't support decoding
if self._use_pointer:
return None
if isinstance(self._ti, MessageType):
return self.decode_case(
f"this->{self.field_name}.emplace_back();\n"
f"value.decode_to_message(this->{self.field_name}.back());"
)
return super().decode_content
content = self._ti.decode_varint
if content is None:
return None
return (
f"case {self.number}: this->{self.field_name}.push_back({content}); break;"
)
@property
def decode_length_content(self) -> str:
# Pointer fields don't support decoding
if self._use_pointer:
return None
content = self._ti.decode_length
if content is None and isinstance(self._ti, MessageType):
# Special handling for non-template message decoding
return f"case {self.number}: this->{self.field_name}.emplace_back(); value.decode_to_message(this->{self.field_name}.back()); break;"
if content is None:
return None
return (
f"case {self.number}: this->{self.field_name}.push_back({content}); break;"
)
@property
def decode_32bit_content(self) -> str:
# Pointer fields don't support decoding
if self._use_pointer:
return None
content = self._ti.decode_32bit
if content is None:
return None
return (
f"case {self.number}: this->{self.field_name}.push_back({content}); break;"
)
@property
def decode_64bit_content(self) -> str:
# Pointer fields don't support decoding
if self._use_pointer:
return None
content = self._ti.decode_64bit
if content is None:
return None
return (
f"case {self.number}: this->{self.field_name}.push_back({content}); break;"
)
@property
def _ti_is_bool(self) -> bool:
@@ -2125,7 +2135,15 @@ class RepeatedTypeInfo(TypeInfo):
return isinstance(self._ti, BoolType)
def _encode_element_call(self, element: str) -> str:
return self._ti.encode_element(self.number, element)
"""Helper to generate encode call for a single element."""
if isinstance(self._ti, EnumType):
return f"ProtoEncode::{self._ti.encode_func}(pos, {self.number}, static_cast<uint32_t>({element}), true);"
# Repeated message elements use encode_sub_message (force=true is default)
if isinstance(self._ti, MessageType):
return f"ProtoEncode::encode_sub_message(pos, buffer, {self.number}, {element});"
return (
f"ProtoEncode::{self._ti.encode_func}(pos, {self.number}, {element}, true);"
)
@property
def encode_content(self) -> str:
@@ -2134,7 +2152,7 @@ class RepeatedTypeInfo(TypeInfo):
# Special handling for const char* elements (when container_no_template contains "const char")
if "const char" in self._container_no_template:
o = f"for (const char *it : *this->{self.field_name}) {{\n"
o += f" {_encode_call(self._ti.encode_func, str(self.number), 'it', 'strlen(it)', force=True)}\n"
o += f" ProtoEncode::{self._ti.encode_func}(pos, {self.number}, it, strlen(it), true);\n"
else:
o = f"for (const auto &it : *this->{self.field_name}) {{\n"
o += f" {self._encode_element_call('it')}\n"
@@ -2513,6 +2531,11 @@ def calculate_message_max_size(desc: descriptor.DescriptorProto) -> int | None:
return total_size
# Contents must never reach the log: dump_to prints only the name
SENSITIVE_MESSAGES = {"NoiseResumeTicket"}
SENSITIVE_MESSAGES_SEEN: set[str] = set()
def build_message_type(
desc: descriptor.DescriptorProto,
base_class_fields: dict[str, list[descriptor.FieldDescriptorProto]],
@@ -2520,7 +2543,10 @@ def build_message_type(
) -> tuple[str, str, str]:
public_content: list[str] = []
protected_content: list[str] = []
decode: list[str] = []
decode_varint: list[str] = []
decode_length: list[str] = []
decode_32bit: list[str] = []
decode_64bit: list[str] = []
encode: list[str] = []
dump: list[str] = []
size_calc: list[str] = []
@@ -2649,8 +2675,22 @@ def build_message_type(
if field.options.HasExtension(pb.field_ifdef):
field_ifdef = field.options.Extensions[pb.field_ifdef]
if case := ti.decode_content:
decode.extend(wrap_with_ifdef(case, field_ifdef))
if ti.decode_varint_content:
decode_varint.extend(
wrap_with_ifdef(ti.decode_varint_content, field_ifdef)
)
if ti.decode_length_content:
decode_length.extend(
wrap_with_ifdef(ti.decode_length_content, field_ifdef)
)
if ti.decode_32bit_content:
decode_32bit.extend(
wrap_with_ifdef(ti.decode_32bit_content, field_ifdef)
)
if ti.decode_64bit_content:
decode_64bit.extend(
wrap_with_ifdef(ti.decode_64bit_content, field_ifdef)
)
if ti.dump_content:
# Check for field_ifdef option for dump as well
field_ifdef = None
@@ -2660,15 +2700,49 @@ def build_message_type(
dump.extend(wrap_with_ifdef(ti.dump_content, field_ifdef))
cpp = ""
if decode:
o = f"void {desc.name}::decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) {{\n"
o += " const ProtoFieldValue value(data, scalar);\n"
o += " switch (tag) {\n"
o += indent("\n".join(decode), " ") + "\n"
if decode_varint:
o = f"bool {desc.name}::decode_varint(uint32_t field_id, proto_varint_value_t value) {{\n"
o += " switch (field_id) {\n"
o += indent("\n".join(decode_varint), " ") + "\n"
o += " default: return false;\n"
o += " }\n"
o += " return true;\n"
o += "}\n"
cpp += o
prot = "void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;"
prot = "bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;"
protected_content.insert(0, prot)
if decode_length:
o = f"bool {desc.name}::decode_length(uint32_t field_id, ProtoLengthDelimited value) {{\n"
o += " switch (field_id) {\n"
o += indent("\n".join(decode_length), " ") + "\n"
o += " default: return false;\n"
o += " }\n"
o += " return true;\n"
o += "}\n"
cpp += o
prot = "bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;"
protected_content.insert(0, prot)
if decode_32bit:
o = f"bool {desc.name}::decode_32bit(uint32_t field_id, Proto32Bit value) {{\n"
o += " switch (field_id) {\n"
o += indent("\n".join(decode_32bit), " ") + "\n"
o += " default: return false;\n"
o += " }\n"
o += " return true;\n"
o += "}\n"
cpp += o
prot = "bool decode_32bit(uint32_t field_id, Proto32Bit value) override;"
protected_content.insert(0, prot)
if decode_64bit:
o = f"bool {desc.name}::decode_64bit(uint32_t field_id, Proto64Bit value) {{\n"
o += " switch (field_id) {\n"
o += indent("\n".join(decode_64bit), " ") + "\n"
o += " default: return false;\n"
o += " }\n"
o += " return true;\n"
o += "}\n"
cpp += o
prot = "bool decode_64bit(uint32_t field_id, Proto64Bit value) override;"
protected_content.insert(0, prot)
# Generate custom decode() override for messages with FixedVector fields
@@ -2715,38 +2789,34 @@ def build_message_type(
)
for line in encode
]
o = f"{speed_attr}uint8_t *{desc.name}::encode_msg(const void *self, ProtoWriteBuffer &buffer PROTO_ENCODE_DEBUG_PARAM) {{\n"
o += f" const auto &msg = *static_cast<const {desc.name} *>(self);\n"
o = f"{speed_attr}uint8_t *{desc.name}::encode(ProtoWriteBuffer &buffer PROTO_ENCODE_DEBUG_PARAM) const {{\n"
o += " uint8_t *__restrict__ pos = buffer.get_pos();\n"
o += indent("\n".join(encode_debug)).replace("this->", "msg.") + "\n"
o += indent("\n".join(encode_debug)) + "\n"
o += " return pos;\n"
o += "}\n"
cpp += o
public_content.append(
"static uint8_t *encode_msg(const void *self, ProtoWriteBuffer &buffer PROTO_ENCODE_DEBUG_PARAM);"
)
public_content.append(
"uint8_t *encode(ProtoWriteBuffer &buffer PROTO_ENCODE_DEBUG_PARAM) const {\n"
" return encode_msg(this, buffer PROTO_ENCODE_DEBUG_ARG);\n"
"}"
prot = (
"uint8_t *encode(ProtoWriteBuffer &buffer PROTO_ENCODE_DEBUG_PARAM) const;"
)
public_content.append(prot)
# If no fields to encode or message doesn't need encoding, the default implementation in ProtoMessage will be used
# Add calculate_size method only if this message needs encoding and has fields
if needs_encode and size_calc and not is_inline_only:
o = f"{speed_attr}uint32_t {desc.name}::calc_size_msg(const void *self) {{\n"
o += f" const auto &msg = *static_cast<const {desc.name} *>(self);\n"
o = f"{speed_attr}uint32_t {desc.name}::calculate_size() const {{\n"
o += " uint32_t size = 0;\n"
o += indent("\n".join(size_calc)).replace("this->", "msg.") + "\n"
o += indent("\n".join(size_calc)) + "\n"
o += " return size;\n"
o += "}\n"
cpp += o
public_content.append("static uint32_t calc_size_msg(const void *self);")
public_content.append(
"uint32_t calculate_size() const { return calc_size_msg(this); }"
)
prot = "uint32_t calculate_size() const;"
public_content.append(prot)
# If no fields to calculate size for or message doesn't need encoding, the default implementation in ProtoMessage will be used
if desc.name in SENSITIVE_MESSAGES:
dump = []
SENSITIVE_MESSAGES_SEEN.add(desc.name)
# dump_to method declaration in header
prot = "#ifdef HAS_PROTO_MESSAGE_DUMP\n"
prot += "const char *dump_to(DumpBuffer &out) const override;\n"
@@ -3654,6 +3724,11 @@ static const char *const TAG = "api.service";
except ImportError:
pass
# A renamed message must fail the build, not silently start dumping secrets
missing = SENSITIVE_MESSAGES - SENSITIVE_MESSAGES_SEEN
if missing:
raise RuntimeError(f"SENSITIVE_MESSAGES not found in api.proto: {missing}")
if __name__ == "__main__":
sys.exit(main())
+25 -1
View File
@@ -14,7 +14,31 @@ top=$(git rev-parse --show-toplevel 2>/dev/null) || exit 0
[ -x "$top/venv/bin/python" ] && exit 0
[ -x "$top/script/setup" ] || exit 0
# Every worktree shares the hooks directory of the checkout it was created
# from, and the script/setup run below is the one from whichever branch was just
# checked out. Older branches install their own pre-commit hook without checking
# for a worktree: that moves the shared hook aside as pre-commit.legacy and
# replaces it with one tied to this worktree's virtual environment, so commits
# break in every checkout. To rule that out, the hooks directory is copied
# before script/setup runs and put back exactly as it was afterwards, including
# removing any file script/setup added.
hooks=$(git rev-parse --path-format=absolute --git-path hooks 2>/dev/null) || exit 0
snap=$(mktemp -d "$hooks/.post-checkout.XXXXXX") || exit 0
cp -p "$hooks"/* "$snap"/ 2>/dev/null
# Clear VIRTUAL_ENV so a checkout made from a shell with an environment already
# activated still gets its own, rather than having the active one repointed at
# this working tree.
exec env -u VIRTUAL_ENV "$top/script/setup"
env -u VIRTUAL_ENV "$top/script/setup"
status=$?
for f in "$hooks"/*; do
[ -e "$snap/${f##*/}" ] || rm -f "$f"
done
# Files are moved rather than copied so a hook that is still running, such as
# this one, is swapped out atomically instead of being rewritten in place.
for f in "$snap"/*; do
cmp -s "$f" "$hooks/${f##*/}" 2>/dev/null || mv -f "$f" "$hooks/${f##*/}"
done
rm -rf "$snap"
exit $status
+17
View File
@@ -1104,6 +1104,10 @@ def get_components_per_integration_fixture() -> dict[str, set[str]]:
_TEST_FUNC_RE = re.compile(r"async def (test_\w+)")
# Any usage form (decorator, pytestmark assignment or list element); only
# test_*.py files are scanned, so the marker docs elsewhere cannot false-hit
_SHARED_YAML_USE_RE = re.compile(r"\bmark\.shared_yaml")
_SHARED_YAML_ARG_RE = re.compile(r"\(\s*[\"'](\w+)[\"']\s*\)")
@cache
@@ -1123,6 +1127,19 @@ def get_fixture_to_test_files() -> dict[str, frozenset[str]]:
for func in _TEST_FUNC_RE.findall(content):
base_name = func.replace("test_", "").partition("[")[0]
result.setdefault(base_name, set()).add(rel_path)
# Shared fixtures are named by marker, not by a test function; each
# decorator must carry a string literal or its fixture would silently
# map to no tests
for use in _SHARED_YAML_USE_RE.finditer(content):
arg = _SHARED_YAML_ARG_RE.match(content, use.end())
if arg is None:
line = content.count("\n", 0, use.start()) + 1
raise ValueError(
f"{rel_path}:{line}: shared_yaml marker must take a "
"single-line string literal so CI test selection can map "
"its fixture"
)
result.setdefault(arg.group(1), set()).add(rel_path)
return {k: frozenset(v) for k, v in result.items()}
+164
View File
@@ -0,0 +1,164 @@
#!/usr/bin/env python3
"""Keep pre-commit hook revs in sync with the requirements files.
Dependabot only bumps the ``package==version`` pins in ``requirements*.txt``.
Some of those tools are pinned a second time as hook ``rev`` values in
``.pre-commit-config.yaml``. This script treats the requirements files as
the source of truth and rewrites the revs to match, editing the config
through yamlrocks so comments and layout survive.
Run without arguments to apply the changes in place, or with ``--check`` to
only report drift (exit status 1 when anything is out of sync).
"""
from __future__ import annotations
import argparse
from dataclasses import dataclass
from pathlib import Path
import re
import sys
from typing import Any
import yamlrocks
REPO_ROOT = Path(__file__).resolve().parent.parent
PRECOMMIT_CONFIG = ".pre-commit-config.yaml"
class SyncError(Exception):
"""A pin could not be located in a requirements file or the config."""
@dataclass(frozen=True)
class SyncTarget:
"""A requirements pin and the pre-commit repo whose rev mirrors it."""
package: str
requirements_file: str
repo: str
SYNC_TARGETS: tuple[SyncTarget, ...] = (
SyncTarget(
"ruff", "requirements_test.txt", "https://github.com/astral-sh/ruff-pre-commit"
),
SyncTarget("flake8", "requirements_test.txt", "https://github.com/PyCQA/flake8"),
SyncTarget(
"pyupgrade", "requirements_test.txt", "https://github.com/asottile/pyupgrade"
),
SyncTarget(
"clang-format",
"requirements_dev.txt",
"https://github.com/pre-commit/mirrors-clang-format",
),
SyncTarget(
"yamllint",
"requirements_dev.txt",
"https://github.com/adrienverge/yamllint.git",
),
)
def read_requirement_version(requirements: str, package: str) -> str | None:
"""Return the ``==`` pin for ``package`` or None when it is not pinned."""
pattern = re.compile(
rf"^{re.escape(package)}==(?P<version>[^\s#]+)",
re.MULTILINE | re.IGNORECASE,
)
match = pattern.search(requirements)
return match.group("version") if match else None
def find_repo_entry(doc: Any, repo: str) -> Any:
"""Return the single ``- repo:`` block for ``repo`` in a pre-commit doc."""
try:
entries = [entry for entry in doc["repos"] if entry["repo"] == repo]
except KeyError as err:
raise SyncError(f"malformed pre-commit config, missing key {err}") from None
if len(entries) != 1:
raise SyncError(
f"expected exactly one block for repo {repo}, found {len(entries)}"
)
return entries[0]
def current_rev(entry: Any, repo: str) -> tuple[str, str]:
"""Split the block's rev into its tag prefix (``v`` or empty) and version."""
if "rev" not in entry:
raise SyncError(f"repo {repo} has no rev")
rev = entry["rev"]
if not isinstance(rev, str):
# A rev such as ``1.0`` parses as a number and cannot be compared or
# rewritten safely; quote it in the config instead.
raise SyncError(f"rev of repo {repo} is not a string: {rev!r}")
prefix = "v" if rev.startswith("v") else ""
return prefix, rev.removeprefix("v")
def sync(root: Path, *, write: bool) -> list[str]:
"""Bring every hook rev in line with its requirements pin.
Returns one description per rev that was (or, when ``write`` is False,
would be) changed. Raises SyncError when a pin cannot be found, which
means SYNC_TARGETS has gone stale and needs updating by hand.
"""
config_path = root / PRECOMMIT_CONFIG
doc = yamlrocks.loads(config_path.read_bytes(), option=yamlrocks.OPT_ROUND_TRIP)
requirements: dict[str, str] = {}
changes: list[str] = []
for target in SYNC_TARGETS:
if target.requirements_file not in requirements:
requirements[target.requirements_file] = (
root / target.requirements_file
).read_text()
version = read_requirement_version(
requirements[target.requirements_file], target.package
)
if version is None:
raise SyncError(
f"{target.requirements_file}: no '{target.package}==' pin found"
)
entry = find_repo_entry(doc, target.repo)
prefix, current = current_rev(entry, target.repo)
if current == version:
continue
changes.append(f"{target.package}: {current} -> {version}")
entry["rev"] = f"{prefix}{version}"
if changes and write:
config_path.write_bytes(doc.to_yaml())
return changes
def main(argv: list[str] | None = None) -> int:
parser = argparse.ArgumentParser(description=__doc__.splitlines()[0])
parser.add_argument(
"--check",
action="store_true",
help="report drift without modifying any file; exit 1 if out of sync",
)
parser.add_argument(
"--root",
type=Path,
default=REPO_ROOT,
help="repository checkout to operate on (default: this checkout)",
)
args = parser.parse_args(argv)
try:
changes = sync(args.root, write=not args.check)
except SyncError as err:
print(f"error: {err}", file=sys.stderr)
return 1
for change in changes:
print(change)
if args.check and changes:
return 1
return 0
if __name__ == "__main__": # pragma: no cover
sys.exit(main())
@@ -249,7 +249,7 @@ static APIBuffer build_infrared_rf_transmit_wire() {
std::memcpy(bytes + len, packed, packed_len);
len += packed_len;
// field 6: modulation = 1 (non-zero so it's actually emitted and exercises
// decode_field for this field, matching the documented layout above).
// decode_varint for this field, matching the documented layout above).
put_byte(0x30);
put_varint(1);
@@ -0,0 +1,32 @@
esphome:
name: test-keyboard-no-label
esp32:
board: esp32dev
framework:
type: esp-idf
spi:
- id: spi_bus
clk_pin: GPIO18
mosi_pin: GPIO23
display:
- platform: mipi_spi
spi_id: spi_bus
model: st7789v
id: tft_display
dimensions:
width: 240
height: 320
cs_pin: GPIO22
dc_pin: GPIO21
auto_clear_enabled: false
invert_colors: false
update_interval: never
lvgl:
displays: tft_display
widgets:
- keyboard:
id: keyboard_widget
@@ -0,0 +1,34 @@
esphome:
name: test-qrcode-no-label
esp32:
board: esp32dev
framework:
type: esp-idf
spi:
- id: spi_bus
clk_pin: GPIO18
mosi_pin: GPIO23
display:
- platform: mipi_spi
spi_id: spi_bus
model: st7789v
id: tft_display
dimensions:
width: 240
height: 320
cs_pin: GPIO22
dc_pin: GPIO21
auto_clear_enabled: false
invert_colors: false
update_interval: never
lvgl:
displays: tft_display
widgets:
- qrcode:
id: qr_widget
size: 100
text: "esphome.io"
@@ -0,0 +1,35 @@
esphome:
name: test-tabview-no-label
esp32:
board: esp32dev
framework:
type: esp-idf
spi:
- id: spi_bus
clk_pin: GPIO18
mosi_pin: GPIO23
display:
- platform: mipi_spi
spi_id: spi_bus
model: st7789v
id: tft_display
dimensions:
width: 240
height: 320
cs_pin: GPIO22
dc_pin: GPIO21
auto_clear_enabled: false
invert_colors: false
update_interval: never
lvgl:
displays: tft_display
widgets:
- tabview:
id: tabview_widget
tabs:
- name: "Tab 1"
id: tab_1
@@ -0,0 +1,32 @@
"""Widgets whose LVGL C implementation creates or references labels
internally (tab titles, key legends, the QR canvas fallback) must declare
the label dependency in ``get_uses()``. Otherwise a config that contains
no ``label`` widget of its own compiles LVGL without ``LV_USE_LABEL`` and
fails at C compile time with undefined ``lv_label_*`` symbols.
"""
from __future__ import annotations
from collections.abc import Callable
from pathlib import Path
import pytest
from esphome.components.lvgl import defines as df
@pytest.mark.parametrize(
"yaml_file",
[
"qrcode_no_label.yaml",
"keyboard_no_label.yaml",
"tabview_no_label.yaml",
],
)
def test_label_less_config_enables_lv_use_label(
generate_main: Callable[[str | Path], str],
component_config_path: Callable[[str], Path],
yaml_file: str,
) -> None:
generate_main(component_config_path(yaml_file))
assert "LV_USE_LABEL" in df.get_defines()
@@ -59,7 +59,7 @@ static void verify_mac(uint64_t mac, size_t expected_bytes) {
#ifdef ESPHOME_DEBUG_API
uint8_t *proto_debug_end_ = api_buf.data() + api_buf.size();
#endif
pos = ProtoEncode::encode_varint_raw_48bit(pos PROTO_ENCODE_DEBUG_ARG, mac);
ProtoEncode::encode_varint_raw_48bit(pos PROTO_ENCODE_DEBUG_ARG, mac);
size_t new_len = pos - api_buf.data();
EXPECT_EQ(new_len, expected_bytes) << "mac=0x" << std::hex << mac << std::dec;
+5
View File
@@ -0,0 +1,5 @@
from tests.testing_helpers import ComponentManifestOverride
def override_manifest(manifest: ComponentManifestOverride) -> None:
manifest.dependencies = manifest.dependencies + ["sensor", "spi"]
@@ -0,0 +1,62 @@
#include <gtest/gtest.h>
#include "esphome/components/atm90e32/atm90e32.h"
namespace esphome::atm90e32::testing {
TEST(ATM90E32OffsetRegisterVerification, AcceptsExactSignedReadback) {
EXPECT_TRUE(offset_register_value_matches(0x007B, 123));
EXPECT_TRUE(offset_register_value_matches(0xFF85, -123));
}
TEST(ATM90E32OffsetRegisterVerification, RejectsMismatchedReadback) {
EXPECT_FALSE(offset_register_value_matches(0x007C, 123));
EXPECT_FALSE(offset_register_value_matches(0xFF84, -123));
}
TEST(ATM90E32OffsetRestoreState, ReportsVerifiedStoredValuesAsRestored) {
const auto state = resolve_offset_restore_state(true, true, false);
EXPECT_TRUE(state.restored);
EXPECT_TRUE(state.values_verified);
}
TEST(ATM90E32OffsetRestoreState, ReportsVerifiedConfigFallbackAsNotRestored) {
const auto state = resolve_offset_restore_state(true, false, true);
EXPECT_FALSE(state.restored);
EXPECT_TRUE(state.values_verified);
}
TEST(ATM90E32OffsetRestoreState, ReportsFailedConfigFallbackAsUnverified) {
const auto state = resolve_offset_restore_state(true, false, false);
EXPECT_FALSE(state.restored);
EXPECT_FALSE(state.values_verified);
}
TEST(ATM90E32OffsetRestoreState, ReportsConfigWithoutStoredValuesAsNotRestored) {
const auto state = resolve_offset_restore_state(false, true, false);
EXPECT_FALSE(state.restored);
EXPECT_TRUE(state.values_verified);
}
TEST(ATM90E32OffsetPersistence, RollsBackStoredValuesOrZeroSentinel) {
const OffsetCalibration previous[3]{{1, -1}, {2, -2}, {3, -3}};
OffsetCalibration rollback[3]{};
prepare_offset_rollback(previous, true, rollback);
for (uint8_t phase = 0; phase < 3; phase++) {
EXPECT_EQ(rollback[phase].first_offset, previous[phase].first_offset);
EXPECT_EQ(rollback[phase].second_offset, previous[phase].second_offset);
}
prepare_offset_rollback(previous, false, rollback);
for (const auto &phase : rollback) {
EXPECT_EQ(phase.first_offset, 0);
EXPECT_EQ(phase.second_offset, 0);
}
}
} // namespace esphome::atm90e32::testing
-43
View File
@@ -59,47 +59,4 @@ TEST(StringRefStartsWith, RefOverloadComparesOnlyTheViewedLength) {
EXPECT_TRUE(ref.starts_with(prefix));
}
// The generated api messages start their encode only string fields as a null pointer with zero
// length; every member must treat that exactly like the default constructed empty string.
TEST(StringRefNullEmpty, BehavesAsEmptyString) {
const StringRef null_empty{nullptr, 0};
const StringRef empty;
EXPECT_TRUE(null_empty.empty());
EXPECT_EQ(null_empty.size(), 0u);
EXPECT_EQ(null_empty.c_str(), nullptr);
EXPECT_TRUE(null_empty == empty);
EXPECT_TRUE(null_empty == "");
EXPECT_TRUE(null_empty == std::string());
EXPECT_EQ(null_empty.compare(empty), 0);
EXPECT_EQ(null_empty.compare(""), 0);
EXPECT_LT(null_empty.compare("a"), 0);
EXPECT_TRUE(null_empty.starts_with(""));
EXPECT_FALSE(null_empty.starts_with("a"));
EXPECT_EQ(null_empty.str(), std::string());
EXPECT_EQ(null_empty.substr(0), std::string());
EXPECT_EQ(null_empty.find('a'), std::string::npos);
EXPECT_EQ(null_empty.find("a"), std::string::npos);
char buf[4] = "xyz";
EXPECT_EQ(null_empty.copy(buf, sizeof(buf)), 0u);
EXPECT_EQ(null_empty.begin(), null_empty.end());
}
TEST(StringRefNullEmpty, ComparesAgainstText) {
const StringRef null_empty{nullptr, 0};
const StringRef text("abc", 3);
EXPECT_FALSE(null_empty == text);
EXPECT_FALSE(text == null_empty);
EXPECT_LT(null_empty.compare(text), 0);
EXPECT_GT(text.compare(null_empty), 0);
EXPECT_TRUE(text.starts_with(null_empty));
}
TEST(StringRefNullEmpty, TwoNullViewsAreEqual) {
const StringRef a{nullptr, 0};
const StringRef b{nullptr, 0};
EXPECT_TRUE(a == b);
EXPECT_EQ(a.compare(b), 0);
EXPECT_TRUE(a.starts_with(b));
}
} // namespace esphome::core::testing
@@ -0,0 +1,224 @@
#include <gtest/gtest.h>
#include <cstring>
#include <noise/protocol.h>
#include "esphome/components/noise/noise.h"
#include "esphome/components/noise/noise_resume.h"
namespace esphome::noise::testing {
// Known-answer vectors shared with the client implementation
// (aioesphomeapi tests/test_noise_resume.py); the two must stay identical
// byte for byte or resumed sessions cannot interoperate.
static const uint8_t KAT_SECRET[RESUME_SECRET_SIZE] = {1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16,
17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30, 31, 32};
static const uint8_t KAT_SESSION_ID[RESUME_SESSION_ID_SIZE] = {0xa0, 0xa1, 0xa2, 0xa3, 0xa4, 0xa5, 0xa6, 0xa7};
static const uint8_t KAT_CLIENT_NONCE[RESUME_NONCE_SIZE] = {0x10, 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17,
0x18, 0x19, 0x1a, 0x1b, 0x1c, 0x1d, 0x1e, 0x1f};
static const uint8_t KAT_SERVER_NONCE[RESUME_NONCE_SIZE] = {0x30, 0x31, 0x32, 0x33, 0x34, 0x35, 0x36, 0x37,
0x38, 0x39, 0x3a, 0x3b, 0x3c, 0x3d, 0x3e, 0x3f};
static const uint8_t KAT_OFFER_MAC[RESUME_MAC_SIZE] = {0xa8, 0x08, 0xea, 0xdb, 0xec, 0x81, 0xa7, 0xcb,
0xf4, 0xca, 0xaa, 0xb8, 0x0d, 0x7f, 0x9d, 0x01};
static const uint8_t KAT_CONFIRM_MAC[RESUME_MAC_SIZE] = {0x09, 0xa3, 0x70, 0x3e, 0xc8, 0x34, 0x77, 0xe9,
0x45, 0xe7, 0xf1, 0x61, 0x9d, 0x4f, 0x6a, 0x76};
static const uint8_t KAT_K_C2D[32] = {0xd6, 0x01, 0xe3, 0xc1, 0x16, 0xa1, 0x64, 0x66, 0xdb, 0xc5, 0x9e,
0xdd, 0x60, 0x2a, 0x64, 0x1e, 0xbe, 0xf5, 0x11, 0x95, 0x98, 0xd2,
0xf2, 0x47, 0x1b, 0xc6, 0x8c, 0x51, 0x8f, 0xbe, 0xb7, 0x23};
static const uint8_t KAT_K_D2C[32] = {0x7f, 0x8d, 0x57, 0x7e, 0x9f, 0xb4, 0xbb, 0xde, 0x86, 0xcd, 0xa9,
0xf4, 0x9b, 0x42, 0xe7, 0x24, 0xc8, 0x49, 0xce, 0x89, 0xd8, 0x96,
0x3f, 0x3c, 0x4b, 0x3f, 0x8f, 0x80, 0xc2, 0x56, 0xab, 0x65};
/// The one place in this file that spells the offer wire layout
static void build_offer(uint8_t *offer, const uint8_t *session_id, const uint8_t *client_nonce, const uint8_t *mac) {
offer[0] = RESUME_OFFER_VERSION;
std::memcpy(offer + RESUME_OFFER_SESSION_ID_OFFSET, session_id, RESUME_SESSION_ID_SIZE);
std::memcpy(offer + RESUME_OFFER_NONCE_OFFSET, client_nonce, RESUME_NONCE_SIZE);
std::memcpy(offer + RESUME_OFFER_MAC_OFFSET, mac, RESUME_MAC_SIZE);
}
/// "NoiseAPIInit" || be16(len) || offer, exactly as the api frame helper mixes it
static constexpr size_t KAT_PROLOGUE_SIZE = 12 + 2 + RESUME_OFFER_SIZE;
static void build_prologue(uint8_t *out, const uint8_t *offer) {
std::memcpy(out, "NoiseAPIInit", 12); // NOLINT(bugprone-not-null-terminated-result)
out[12] = 0x00;
out[13] = RESUME_OFFER_SIZE;
std::memcpy(out + 14, offer, RESUME_OFFER_SIZE);
}
static void build_offer_for_ticket(uint8_t *offer, const ResumeTicket &ticket, const uint8_t *client_nonce) {
uint8_t mac[RESUME_MAC_SIZE];
ASSERT_TRUE(resume_compute_offer_mac(ticket.secret, ticket.session_id, client_nonce, mac));
build_offer(offer, ticket.session_id, client_nonce, mac);
}
/// Test access to the protected slots so a test can plant the KAT ticket
struct TestCache : ResumeTicketCache {
void plant(const uint8_t *session_id, const uint8_t *secret) {
std::memcpy(this->slots_[0].session_id, session_id, RESUME_SESSION_ID_SIZE);
std::memcpy(this->slots_[0].secret, secret, RESUME_SECRET_SIZE);
this->used_mask_ |= 1u;
}
};
TEST(NoiseResumeKat, ConfirmMacMatchesClientImplementation) {
uint8_t mac[RESUME_MAC_SIZE];
ASSERT_TRUE(resume_compute_confirm_mac(KAT_SECRET, KAT_CLIENT_NONCE, KAT_SERVER_NONCE, mac));
EXPECT_EQ(std::memcmp(mac, KAT_CONFIRM_MAC, RESUME_MAC_SIZE), 0);
}
TEST(NoiseResumeKat, OfferMacMatchesClientImplementation) {
uint8_t mac[RESUME_MAC_SIZE];
ASSERT_TRUE(resume_compute_offer_mac(KAT_SECRET, KAT_SESSION_ID, KAT_CLIENT_NONCE, mac));
EXPECT_EQ(std::memcmp(mac, KAT_OFFER_MAC, RESUME_MAC_SIZE), 0);
}
TEST(NoiseResumeKat, KeyDerivationMatchesClientImplementation) {
// Prologue used by the shared vectors: "NoiseAPIInit" + be16(41) + a
// 41-byte offer whose MAC field is 16 bytes of 0xEE
uint8_t mac_filler[RESUME_MAC_SIZE];
std::memset(mac_filler, 0xEE, sizeof(mac_filler));
uint8_t offer[RESUME_OFFER_SIZE];
build_offer(offer, KAT_SESSION_ID, KAT_CLIENT_NONCE, mac_filler);
uint8_t prologue[KAT_PROLOGUE_SIZE];
build_prologue(prologue, offer);
uint8_t k_c2d[32], k_d2c[32];
ASSERT_TRUE(
resume_derive_keys(KAT_SECRET, KAT_CLIENT_NONCE, KAT_SERVER_NONCE, prologue, sizeof(prologue), k_c2d, k_d2c));
EXPECT_EQ(std::memcmp(k_c2d, KAT_K_C2D, 32), 0);
EXPECT_EQ(std::memcmp(k_d2c, KAT_K_D2C, 32), 0);
}
TEST(NoiseResumeCache, TryAcceptConsumesTicketOnceAndProvesPossession) {
TestCache cache;
cache.plant(KAT_SESSION_ID, KAT_SECRET);
uint8_t offer[RESUME_OFFER_SIZE];
build_offer(offer, KAT_SESSION_ID, KAT_CLIENT_NONCE, KAT_OFFER_MAC);
uint8_t prologue[KAT_PROLOGUE_SIZE];
build_prologue(prologue, offer);
uint8_t ext[RESUME_ACCEPT_SIZE];
NoiseCipherState *send = nullptr, *recv = nullptr;
ASSERT_EQ(cache.try_accept(offer, sizeof(offer), prologue, sizeof(prologue), ext, sizeof(ext), send, recv),
RESUME_ACCEPT_SIZE);
ASSERT_NE(send, nullptr);
ASSERT_NE(recv, nullptr);
// The extension proves possession: verify like the client does
EXPECT_EQ(ext[0], RESUME_ACCEPT_VERSION);
const uint8_t *server_nonce = ext + 1;
uint8_t expected_confirm[RESUME_MAC_SIZE];
ASSERT_TRUE(resume_compute_confirm_mac(KAT_SECRET, KAT_CLIENT_NONCE, server_nonce, expected_confirm));
EXPECT_EQ(std::memcmp(ext + 1 + RESUME_NONCE_SIZE, expected_confirm, RESUME_MAC_SIZE), 0);
// The ciphers must interoperate with the documented key derivation
uint8_t k_c2d[32], k_d2c[32];
ASSERT_TRUE(resume_derive_keys(KAT_SECRET, KAT_CLIENT_NONCE, server_nonce, prologue, sizeof(prologue), k_c2d, k_d2c));
NoiseCipherState *client_send = resume_make_cipher(k_c2d);
ASSERT_NE(client_send, nullptr);
uint8_t buf[64] = "resumed";
NoiseBuffer nb;
noise_buffer_init(nb);
noise_buffer_set_inout(nb, buf, 7, sizeof(buf));
ASSERT_EQ(noise_cipherstate_encrypt(client_send, &nb), NOISE_ERROR_NONE);
ASSERT_EQ(noise_cipherstate_decrypt(recv, &nb), NOISE_ERROR_NONE);
EXPECT_EQ(std::memcmp(buf, "resumed", 7), 0);
noise_cipherstate_free(client_send);
noise_cipherstate_free(send);
noise_cipherstate_free(recv);
// Single use: the same offer must miss the second time
NoiseCipherState *send2 = nullptr, *recv2 = nullptr;
EXPECT_EQ(cache.try_accept(offer, sizeof(offer), prologue, sizeof(prologue), ext, sizeof(ext), send2, recv2), 0u);
EXPECT_EQ(send2, nullptr);
EXPECT_EQ(recv2, nullptr);
}
TEST(NoiseResumeCache, BadMacOrMalformedOfferLeavesTicketIntact) {
TestCache cache;
cache.plant(KAT_SESSION_ID, KAT_SECRET);
uint8_t offer[RESUME_OFFER_SIZE];
uint8_t bad_mac[RESUME_MAC_SIZE];
std::memcpy(bad_mac, KAT_OFFER_MAC, RESUME_MAC_SIZE);
bad_mac[0] ^= 0x01;
build_offer(offer, KAT_SESSION_ID, KAT_CLIENT_NONCE, bad_mac);
uint8_t prologue[1] = {0};
uint8_t ext[RESUME_ACCEPT_SIZE];
NoiseCipherState *send = nullptr, *recv = nullptr;
// A forged offer must not burn the ticket
EXPECT_EQ(cache.try_accept(offer, sizeof(offer), prologue, sizeof(prologue), ext, sizeof(ext), send, recv), 0u);
// Wrong size or version must be recognized as "no offer"
build_offer(offer, KAT_SESSION_ID, KAT_CLIENT_NONCE, KAT_OFFER_MAC);
EXPECT_EQ(cache.try_accept(offer, sizeof(offer) - 1, prologue, sizeof(prologue), ext, sizeof(ext), send, recv), 0u);
offer[0] = 0x7f;
EXPECT_EQ(cache.try_accept(offer, sizeof(offer), prologue, sizeof(prologue), ext, sizeof(ext), send, recv), 0u);
offer[0] = RESUME_OFFER_VERSION;
// No room for the extension must also decline without burning it
EXPECT_EQ(cache.try_accept(offer, sizeof(offer), prologue, sizeof(prologue), ext, sizeof(ext) - 1, send, recv), 0u);
// The genuine offer still redeems
EXPECT_EQ(cache.try_accept(offer, sizeof(offer), prologue, sizeof(prologue), ext, sizeof(ext), send, recv),
RESUME_ACCEPT_SIZE);
noise_cipherstate_free(send);
noise_cipherstate_free(recv);
}
TEST(NoiseResumeCache, SetPskForgetsTickets) {
NoiseContext ctx;
ResumeTicket ticket;
ASSERT_TRUE(ctx.resume_cache().issue(ticket));
psk_t psk{};
psk[0] = 1;
ctx.set_psk(psk.data());
uint8_t offer[RESUME_OFFER_SIZE];
build_offer_for_ticket(offer, ticket, KAT_CLIENT_NONCE);
uint8_t prologue[KAT_PROLOGUE_SIZE];
build_prologue(prologue, offer);
uint8_t ext[RESUME_ACCEPT_SIZE];
NoiseCipherState *send = nullptr, *recv = nullptr;
EXPECT_EQ(
ctx.resume_cache().try_accept(offer, sizeof(offer), prologue, sizeof(prologue), ext, sizeof(ext), send, recv),
0u);
EXPECT_EQ(send, nullptr);
EXPECT_EQ(recv, nullptr);
}
TEST(NoiseResumeCache, IssueRotatesSlotsAndClearForgetsAll) {
ResumeTicketCache cache;
ResumeTicket tickets[ResumeTicketCache::SLOTS + 1];
for (auto &ticket : tickets) {
ASSERT_TRUE(cache.issue(ticket));
}
uint8_t offer[RESUME_OFFER_SIZE];
uint8_t prologue[1] = {0};
uint8_t ext[RESUME_ACCEPT_SIZE];
// The oldest ticket was evicted by the one-past-capacity issue
build_offer_for_ticket(offer, tickets[0], KAT_CLIENT_NONCE);
NoiseCipherState *send = nullptr, *recv = nullptr;
EXPECT_EQ(cache.try_accept(offer, sizeof(offer), prologue, sizeof(prologue), ext, sizeof(ext), send, recv), 0u);
// The rest remain redeemable
for (int i = 1; i <= ResumeTicketCache::SLOTS; i++) {
build_offer_for_ticket(offer, tickets[i], KAT_CLIENT_NONCE);
EXPECT_EQ(cache.try_accept(offer, sizeof(offer), prologue, sizeof(prologue), ext, sizeof(ext), send, recv),
RESUME_ACCEPT_SIZE);
noise_cipherstate_free(send);
noise_cipherstate_free(recv);
send = recv = nullptr;
}
// clear() forgets everything
ResumeTicket ticket;
ASSERT_TRUE(cache.issue(ticket));
cache.clear();
build_offer_for_ticket(offer, ticket, KAT_CLIENT_NONCE);
EXPECT_EQ(cache.try_accept(offer, sizeof(offer), prologue, sizeof(prologue), ext, sizeof(ext), send, recv), 0u);
}
} // namespace esphome::noise::testing
@@ -0,0 +1,29 @@
# Tuya without any network component (no wifi/ethernet/api), as used on
# serial-only or BLE-only Tuya MCU boards. Regression test for
# https://github.com/esphome/esphome/issues/18942
substitutions:
status_pin: P6
packages:
uart: !include ../../test_build_components/common/uart/bk72xx-ard.yaml
tuya:
status_pin: ${status_pin}
binary_sensor:
- platform: tuya
id: tuya_presence
sensor_datapoint: 101
sensor:
- platform: tuya
id: tuya_light_intensity
sensor_datapoint: 103
number:
- platform: tuya
id: tuya_far_detection
number_datapoint: 109
min_value: 0
max_value: 600
step: 1
+7
View File
@@ -21,6 +21,13 @@ The `yaml_config` fixture automatically loads YAML configurations based on the t
- The fixture file must exist or the test will fail with a clear error message
- The fixture automatically injects a dynamic port number into the API configuration
Tests marked `@pytest.mark.shared_yaml("name")` load `fixtures/name.yaml` instead
of the test-named file and compile it in a shared, hash-keyed build directory, so
the whole group pays one full compile and each test only a relink. The marker
argument must be a single-line string literal (CI test selection maps fixtures to
test files by scanning for it), and marked tests must hand the `yaml_config`
content to `run_compiled` unmodified.
### Key Fixtures
- `run_compiled` - Combines write, compile, and run operations into a single context manager
+335 -81
View File
@@ -4,17 +4,22 @@ from __future__ import annotations
import asyncio
from collections.abc import AsyncGenerator, Callable, Generator
from contextlib import AbstractAsyncContextManager, asynccontextmanager
from contextlib import AbstractAsyncContextManager, asynccontextmanager, suppress
import fcntl
from functools import cache
import hashlib
import logging
import os
from pathlib import Path
import platform
import re
import shutil
import signal
import socket
import subprocess
import sys
import tempfile
import time
from typing import TextIO
from aioesphomeapi import APIClient, APIConnectionError, LogParser, ReconnectLogic
@@ -23,7 +28,13 @@ import pytest_asyncio
import esphome.config
from esphome.core import CORE
from esphome.helpers import get_usable_cpu_count
from esphome.helpers import (
get_usable_cpu_count,
read_file,
rmtree,
write_file,
write_file_if_changed,
)
from esphome.platformio.toolchain import get_idedata
from .const import (
@@ -56,6 +67,21 @@ import pty # not available on Windows
pytest.register_assert_rewrite("tests.integration.entity_utils")
def pytest_configure(config: pytest.Config) -> None:
config.addinivalue_line(
"markers",
"shared_yaml(name): load fixtures/<name>.yaml and compile it in a shared, "
"hash-keyed incremental build directory",
)
FIXTURES_DIR = Path(__file__).parent / "fixtures"
REPO_ROOT = Path(__file__).resolve().parent.parent.parent
# CI caches parts of this path; keep in sync with ci.yml integration-tests.
INTEGRATION_TESTS_ROOT = Path.home() / ".esphome-integration-tests"
def _get_platformio_env(cache_dir: Path) -> dict[str, str]:
"""Get environment variables for PlatformIO with shared cache."""
env = os.environ.copy()
@@ -78,7 +104,7 @@ def _get_platformio_env(cache_dir: Path) -> dict[str, str]:
)
# Compile with THIS tree's esphome sources, not wherever the venv's editable
# install points (which may be a different git worktree or checkout).
repo_root = str(Path(__file__).resolve().parent.parent.parent)
repo_root = str(REPO_ROOT)
existing = env.get("PYTHONPATH")
env["PYTHONPATH"] = f"{repo_root}{os.pathsep}{existing}" if existing else repo_root
return env
@@ -88,8 +114,7 @@ def _get_platformio_env(cache_dir: Path) -> dict[str, str]:
def shared_platformio_cache() -> Generator[Path]:
"""Initialize a shared PlatformIO cache for all integration tests."""
# Use a dedicated directory for integration tests to avoid conflicts.
# CI caches parts of this path; keep in sync with ci.yml integration-tests.
test_cache_dir = Path.home() / ".esphome-integration-tests"
test_cache_dir = INTEGRATION_TESTS_ROOT
cache_dir = test_cache_dir / "platformio"
# Use a lock file in the home directory to ensure only one process initializes the cache
@@ -112,7 +137,9 @@ def shared_platformio_cache() -> Generator[Path]:
init_dir = Path(tmpdir)
fixture_path = Path(__file__).parent / "fixtures" / "cache_init.yaml"
config_path = init_dir / "cache_init.yaml"
config_path.write_text(fixture_path.read_text())
config_path.write_text(
fixture_path.read_text(encoding="utf-8"), encoding="utf-8"
)
# Run compilation to populate the cache
# We must succeed here to avoid race conditions where multiple
@@ -162,13 +189,6 @@ def integration_test_dir() -> Generator[Path]:
yield Path(tmpdir)
@pytest.fixture
def isolated_preferences(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None:
"""Host preferences persist per device name; give the test its own so a
provisioned key never leaks into another run."""
monkeypatch.setenv("ESPHOME_PREFDIR", str(tmp_path / "prefs"))
@pytest.fixture
def reserved_tcp_port() -> Generator[tuple[int, socket.socket]]:
"""Reserve an unused TCP port by holding the socket open."""
@@ -188,21 +208,29 @@ def unused_tcp_port(reserved_tcp_port: tuple[int, socket.socket]) -> int:
return reserved_tcp_port[0]
@pytest.fixture(autouse=True)
def isolated_preferences(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> Path:
"""Give every test its own host prefs dir; prefs are keyed only by device
name, which tests sharing a fixture also share."""
prefdir = tmp_path / "prefs"
monkeypatch.setenv("ESPHOME_PREFDIR", str(prefdir))
return prefdir
@pytest_asyncio.fixture
async def yaml_config(request: pytest.FixtureRequest, unused_tcp_port: int) -> str:
"""Load YAML configuration based on test name."""
# Get the test function name
test_name: str = request.node.name
# Extract the base test name (remove test_ prefix and any parametrization)
base_name = test_name.replace("test_", "").partition("[")[0]
shared_name = _shared_yaml_name(request)
# Base test name: test_ prefix and any parametrization stripped
base_name = shared_name or request.node.name.replace("test_", "").partition("[")[0]
# Load the fixture file
fixture_path = Path(__file__).parent / "fixtures" / f"{base_name}.yaml"
fixture_path = FIXTURES_DIR / f"{base_name}.yaml"
if not fixture_path.exists():
raise FileNotFoundError(f"Fixture file not found: {fixture_path}")
loop = asyncio.get_running_loop()
content = await loop.run_in_executor(None, fixture_path.read_text)
content = await loop.run_in_executor(None, read_file, fixture_path)
# Replace the port in the config if it contains api section
if "api:" in content:
@@ -226,11 +254,13 @@ async def yaml_config(request: pytest.FixtureRequest, unused_tcp_port: int) -> s
# Replace external component path placeholder if present
if "EXTERNAL_COMPONENT_PATH" in content:
external_components_path = str(
Path(__file__).parent / "fixtures" / "external_components"
)
external_components_path = str(FIXTURES_DIR / "external_components")
content = content.replace("EXTERNAL_COMPONENT_PATH", external_components_path)
if shared_name is not None:
# _compile verifies the marked test compiles this content unmodified
request.node._shared_yaml_content = content
return content
@@ -240,24 +270,218 @@ async def write_yaml_config(
) -> AsyncGenerator[ConfigWriter]:
"""Write YAML configuration to a file."""
# Get the test name for default filename
test_name = request.node.name
base_name = test_name.replace("test_", "").split("[")[0]
base_name = request.node.name.replace("test_", "").partition("[")[0]
async def _write_config(content: str, filename: str | None = None) -> Path:
if filename is None:
filename = f"{base_name}.yaml"
config_path = integration_test_dir / filename
loop = asyncio.get_running_loop()
await loop.run_in_executor(None, config_path.write_text, content)
await loop.run_in_executor(None, write_file, config_path, content)
return config_path
yield _write_config
# Deliberately not CI-cached (ci.yml caches only platformio/ subpaths); stale
# dirs for a fixture are pruned when its content hash changes.
SHARED_BUILDS_ROOT = INTEGRATION_TESTS_ROOT / "builds"
# In the dir name (not just the hash) so pruning stays inside this checkout
_REPO_KEY = hashlib.sha256(str(REPO_ROOT).encode()).hexdigest()[:8]
# Give a contended shared build lock time for a full cold compile ahead of us
_SHARED_LOCK_TIMEOUT_S = 900
_SHARED_LOCK_POLL_S = 0.1
_SHARED_LOCK_REPORT_S = 30
# Reclaims dirs orphaned by fixture renames or deleted checkouts
_STALE_BUILD_MAX_AGE_S = 30 * 24 * 3600
# ELF path per shared build dir; constant once compiled, so resolve it only once
_shared_elf_paths: dict[Path, Path] = {}
# Dirs this process already swept; pruning is session-scoped work
_pruned_dirs: set[Path] = set()
def _shared_yaml_name(request: pytest.FixtureRequest) -> str | None:
"""Name passed to the shared_yaml marker, or None when unmarked."""
marker = request.node.get_closest_marker("shared_yaml")
if marker is None:
return None
# Exactly one \w+ positional arg: the name doubles as a build dir
# component, and CI test selection (script/helpers.py) parses the same shape
if (
len(marker.args) != 1
or marker.kwargs
or not re.fullmatch(r"\w+", str(marker.args[0]))
):
raise ValueError(
"shared_yaml marker requires exactly one \\w+ fixture name literal"
)
return marker.args[0]
def _shared_build_prefix(name: str) -> str:
return f"{name}-{_REPO_KEY}-"
@cache
def _shared_build_dir(name: str) -> Path:
"""Dir keyed by checkout and fixture source, before per-test injections."""
key = hashlib.sha256((FIXTURES_DIR / f"{name}.yaml").read_bytes()).hexdigest()[:16]
return SHARED_BUILDS_ROOT / (_shared_build_prefix(name) + key)
def _read_stamp(stamp: Path, shared_dir: Path) -> Path | None:
"""ELF path recorded by the last completed compile, or None."""
try:
text = stamp.read_text(encoding="utf-8").strip()
except FileNotFoundError:
return None
except OSError as err:
print(f"Cannot read {stamp}: {err}")
return None
if not text:
print(f"Ignoring empty stamp {stamp}")
return None
built = Path(text)
# Never trust a stamp pointing outside its own build dir as an unlink target
if shared_dir.resolve() in built.resolve().parents:
return built
print(f"Ignoring stamp {stamp} pointing outside {shared_dir}")
return None
def _unused_since(stale: Path, cutoff: float) -> bool:
"""Whether a build dir looks untouched since cutoff; unknown counts as used."""
# Newest of the .built stamp (rewritten by every completed compile) and the
# dir itself (freshened by a worker claiming the dir before locking)
newest: float | None = None
for probe in (stale / ".built", stale):
try:
mtime = probe.stat().st_mtime
except FileNotFoundError:
continue
except NotADirectoryError:
return True # a stray file where a dir should be; reclaimable
except OSError as err:
print(f"Cannot age-probe {stale}: {err}")
return False # unknown never authorizes deletion
newest = mtime if newest is None else max(newest, mtime)
return newest is not None and newest < cutoff
def _prune_stale_builds(name: str, keep: Path) -> None:
"""Remove outdated build dirs (blocking, run in executor): this checkout's
other dirs for the fixture, plus anything untouched for 30 days. Tolerates
other workers pruning the same dirs concurrently."""
cutoff = time.time() - _STALE_BUILD_MAX_AGE_S
prefix = _shared_build_prefix(name)
for stale in SHARED_BUILDS_ROOT.iterdir():
if stale == keep:
continue
same_fixture = stale.name.startswith(prefix)
if not same_fixture and not _unused_since(stale, cutoff):
continue
# Creating .lock bumps the dir mtime, so remember whether the re-probe
# under the lock can trust it
lock_preexisting = (stale / ".lock").exists()
try:
lock_file = (stale / ".lock").open("w")
except FileNotFoundError:
continue # pruned by another worker meanwhile
except NotADirectoryError:
print(f"Removing stray file {stale}")
stale.unlink(missing_ok=True)
continue
except OSError as err:
print(f"Cannot prune {stale}: {err}")
continue
with lock_file:
try:
fcntl.flock(lock_file.fileno(), fcntl.LOCK_EX | fcntl.LOCK_NB)
except BlockingIOError:
continue # still in use by another run
# Re-probe under the lock: a worker freshens its dir before
# locking, so a just-claimed dir no longer looks unused. A dir
# whose .lock we just created cannot be held by anyone, and our
# own open bumped its mtime, so its pre-open probe stands
if (
lock_preexisting
and not same_fixture
and not _unused_since(stale, cutoff)
):
continue
# rmtree tolerates races; a leftover partial tree only costs a
# rebuild, since the ELF is deleted before every compile
try:
rmtree(stale)
except OSError as err:
print(f"Failed to prune {stale}: {err}")
async def _run_esphome_compile(
config_path: Path, cwd: Path, env: dict[str, str]
) -> None:
"""Run `esphome compile`, retrying up to 3 times on a segfault."""
max_retries = 3
for attempt in range(max_retries):
# Compile using subprocess, inheriting stdout/stderr to show progress
proc = await asyncio.create_subprocess_exec(
sys.executable,
"-m",
"esphome",
"compile",
str(config_path),
cwd=cwd,
stdout=None, # Inherit stdout
stderr=None, # Inherit stderr
stdin=asyncio.subprocess.DEVNULL,
# Start in a new process group to isolate signal handling
start_new_session=True,
env=env,
close_fds=False,
)
await proc.wait()
if proc.returncode == 0:
break
if proc.returncode == -11 and attempt < max_retries - 1:
# Segfault (-11 = SIGSEGV), retry
print(
f"Compilation segfaulted (attempt {attempt + 1}/{max_retries}), retrying..."
)
await asyncio.sleep(1) # Brief pause before retry
continue
raise RuntimeError(
f"Failed to compile {config_path}, return code: {proc.returncode}. "
f"Run with 'pytest -s' to see compilation output."
)
def _resolve_compiled_binary(config_path: Path) -> Path:
"""Load the config to learn the compiled ELF path (blocking, run in executor)."""
CORE.reset() # Reset CORE state between test runs
CORE.config_path = config_path
config = esphome.config.read_config(
{"command": "compile", "config": str(config_path)}
)
if config is None:
raise RuntimeError(f"Failed to read config from {config_path}")
idedata = get_idedata(config)
binary_path = Path(idedata.firmware_elf_path)
if not binary_path.exists():
raise RuntimeError(f"Compiled binary not found at {binary_path}")
return binary_path
@pytest_asyncio.fixture
async def compile_esphome(
integration_test_dir: Path,
shared_platformio_cache: Path,
request: pytest.FixtureRequest,
) -> AsyncGenerator[CompileFunction]:
"""Compile an ESPHome configuration and return the binary path."""
@@ -265,66 +489,96 @@ async def compile_esphome(
# Use the shared PlatformIO cache for faster compilation
# This avoids re-downloading dependencies for each test
env = _get_platformio_env(shared_platformio_cache)
# Retry compilation up to 3 times if we get a segfault
max_retries = 3
for attempt in range(max_retries):
# Compile using subprocess, inheriting stdout/stderr to show progress
proc = await asyncio.create_subprocess_exec(
sys.executable,
"-m",
"esphome",
"compile",
str(config_path),
cwd=integration_test_dir,
stdout=None, # Inherit stdout
stderr=None, # Inherit stderr
stdin=asyncio.subprocess.DEVNULL,
# Start in a new process group to isolate signal handling
start_new_session=True,
env=env,
close_fds=False,
)
await proc.wait()
if proc.returncode == 0:
# Success!
break
if proc.returncode == -11 and attempt < max_retries - 1:
# Segfault (-11 = SIGSEGV), retry
print(
f"Compilation segfaulted (attempt {attempt + 1}/{max_retries}), retrying..."
)
await asyncio.sleep(1) # Brief pause before retry
continue
# Other error or final retry
raise RuntimeError(
f"Failed to compile {config_path}, return code: {proc.returncode}. "
f"Run with 'pytest -s' to see compilation output."
)
# Load the config to get idedata (blocking call, must use executor)
loop = asyncio.get_running_loop()
def _read_config_and_get_binary():
CORE.reset() # Reset CORE state between test runs
CORE.config_path = config_path
config = esphome.config.read_config(
{"command": "compile", "config": str(config_path)}
name = _shared_yaml_name(request)
if name is None:
await _run_esphome_compile(config_path, integration_test_dir, env)
return await loop.run_in_executor(
None, _resolve_compiled_binary, config_path
)
if config is None:
raise RuntimeError(f"Failed to read config from {config_path}")
# Get the compiled binary path
idedata = get_idedata(config)
return Path(idedata.firmware_elf_path)
binary_path = await loop.run_in_executor(None, _read_config_and_get_binary)
if not binary_path.exists():
raise RuntimeError(f"Compiled binary not found at {binary_path}")
return binary_path
# Shared fixture: build in a hash-keyed dir so tests sharing a config
# pay one full compile and later only a main.cpp (port) rebuild + relink
shared_dir = _shared_build_dir(name)
shared_dir.mkdir(parents=True, exist_ok=True)
# Freshen the dir before locking so a concurrent age sweep, which
# re-probes under the lock, never reaps a dir a worker just claimed;
# if a peer reaped it already, the guarded lock open recreates it
with suppress(FileNotFoundError):
os.utime(shared_dir)
if shared_dir not in _pruned_dirs:
_pruned_dirs.add(shared_dir)
await loop.run_in_executor(None, _prune_stale_builds, name, shared_dir)
shared_config = shared_dir / f"{name}.yaml"
private_binary = integration_test_dir / f"{name}.elf"
content = await loop.run_in_executor(None, read_file, config_path)
if content != getattr(request.node, "_shared_yaml_content", None):
# The dir is keyed by the fixture source; a mutated config would be
# cached under a hash that does not describe it
raise RuntimeError(
"shared_yaml tests must compile the yaml_config content unmodified"
)
# flock serializes concurrent xdist workers; closing the fd releases it.
# Hand-rolled rather than filelock.FileLock: non-blocking retries keep
# the wait cancellable, while a blocking acquire in an executor thread
# would survive test cancellation holding the fd
try:
lock_file = (shared_dir / ".lock").open("w")
except FileNotFoundError:
# A peer run pruning divergent hashes reaped the dir between our
# mkdir and this open; recreate it and pay a full rebuild
shared_dir.mkdir(parents=True, exist_ok=True)
lock_file = (shared_dir / ".lock").open("w")
with lock_file:
start = time.monotonic()
last_report = start
while True:
try:
fcntl.flock(lock_file.fileno(), fcntl.LOCK_EX | fcntl.LOCK_NB)
break
except BlockingIOError:
now = time.monotonic()
if now - start > _SHARED_LOCK_TIMEOUT_S:
raise RuntimeError(
f"Timed out waiting for the {shared_dir} lock"
) from None
if now - last_report >= _SHARED_LOCK_REPORT_S:
last_report = now
print(
f"Waited {now - start:.0f}s for another worker's "
f"build of {shared_dir.name}"
)
await asyncio.sleep(_SHARED_LOCK_POLL_S)
# .built carries the ELF path of the last completed compile, so
# later workers skip the config re-read in _resolve_compiled_binary
stamp = shared_dir / ".built"
if (built := _shared_elf_paths.get(shared_dir)) is None:
built = await loop.run_in_executor(None, _read_stamp, stamp, shared_dir)
# Delete the ELF before compiling: whatever exists afterwards is
# this compile's output, so no staleness check is ever needed.
# With no usable stamp, sweep any leftover at the known layout
if built is not None:
built.unlink(missing_ok=True)
else:
# Layout-agnostic: ESPHOME_BUILD_PATH can move the build tree
for leftover in shared_dir.rglob("program"):
if leftover.is_file():
leftover.unlink()
await loop.run_in_executor(
None, write_file_if_changed, shared_config, content
)
await _run_esphome_compile(shared_config, shared_dir, env)
if built is None or not built.exists():
built = await loop.run_in_executor(
None, _resolve_compiled_binary, shared_config
)
_shared_elf_paths[shared_dir] = built
await loop.run_in_executor(None, write_file, stamp, str(built))
# Copy out before unlocking: another worker may relink firmware.elf
# while this test is still running its private copy
await loop.run_in_executor(None, shutil.copy2, built, private_binary)
return private_binary
yield _compile
@@ -1,43 +0,0 @@
esphome:
name: api-decode-wire-types-test
host:
api:
logger:
level: DEBUG
switch:
- platform: template
name: "Wire Switch"
optimistic: true
output:
- platform: template
id: wire_dim
type: float
write_action:
- lambda: ""
light:
- platform: monochromatic
name: "Wire Light"
output: wire_dim
default_transition_length: 0s
effects:
- pulse:
name: Pulse
text:
- platform: template
name: "Wire Text"
optimistic: true
mode: text
min_length: 0
max_length: 255
number:
- platform: template
name: "Wire Number"
optimistic: true
min_value: -1000
max_value: 1000
step: 0.5
@@ -1,11 +0,0 @@
esphome:
name: api-empty-message-test
host:
api:
logger:
level: DEBUG
switch:
- platform: template
name: "Empty Message Switch"
optimistic: true
@@ -1,58 +0,0 @@
esphome:
name: api-encode-boundaries-test
# Top-level area fills DeviceInfoResponse.suggested_area (field 16, a two-byte tag)
area:
id: kitchen_area
name: Kitchen
on_boot:
- sensor.template.publish:
id: zero_then_value
state: 0.0
host:
api:
logger:
level: DEBUG
sensor:
- platform: template
name: "Zero Then Value"
id: zero_then_value
# Negative int32 takes the ten byte varint path
accuracy_decimals: -2
update_interval: never
text_sensor:
- platform: template
name: "Long Text"
id: long_text
update_interval: never
number:
- platform: template
name: "Negative Number"
optimistic: true
min_value: -1000
max_value: 1000
step: 0.5
initial_value: -123.5
select:
- platform: template
name: "Long Option Select"
optimistic: true
options:
- short
- "option-with-a-name-long-enough-that-its-length-prefix-needs-two-varint-bytes-when-the-list-entities-response-is-encoded-xxxxxxxxxx"
initial_option: short
button:
- platform: template
name: "Publish Values"
on_press:
- sensor.template.publish:
id: zero_then_value
state: 12.5
- text_sensor.template.publish:
id: long_text
state: !lambda return std::string(200, 'y');
@@ -0,0 +1,9 @@
esphome:
name: host-noise-resume
host:
api:
encryption:
key: N4Yle5YirwZhPiHHsdZLdOA73ndj/84veVaLhTvxCuU=
# VERY_VERBOSE so the frame helper logs "Session resumed!"
logger:
level: VERY_VERBOSE
@@ -1,58 +0,0 @@
esphome:
name: test-batch-window-filters
host:
api:
batch_delay: 0ms # Disable batching to receive all state updates
logger:
level: DEBUG
# Template sensor that we'll use to publish values
sensor:
- platform: template
name: "Source Sensor"
id: source_sensor
accuracy_decimals: 2
# Batch window filters (window_size == send_every) - use streaming filters
- platform: copy
source_id: source_sensor
name: "Min Sensor"
id: min_sensor
filters:
- min:
window_size: 5
send_every: 5
send_first_at: 1
- platform: copy
source_id: source_sensor
name: "Max Sensor"
id: max_sensor
filters:
- max:
window_size: 5
send_every: 5
send_first_at: 1
- platform: copy
source_id: source_sensor
name: "Moving Avg Sensor"
id: moving_avg_sensor
filters:
- sliding_window_moving_average:
window_size: 5
send_every: 5
send_first_at: 1
# Button to trigger publishing test values
button:
- platform: template
name: "Publish Values Button"
id: publish_button
on_press:
- lambda: |-
// Publish 10 values: 1.0, 2.0, ..., 10.0
for (int i = 1; i <= 10; i++) {
id(source_sensor).publish_state(float(i));
}
@@ -1,111 +0,0 @@
esphome:
name: uart-mock-modbus-cli-rw
host:
api:
logger:
level: VERBOSE
external_components:
- source:
type: local
path: EXTERNAL_COMPONENT_PATH
# Dummy uart entry to satisfy modbus's DEPENDENCIES = ["uart"]
# The actual UART bus used is the uart_mock component below
uart:
baud_rate: 115200
port: /dev/null
# Two virtual buses looped back to each other: the client's transmissions reach the server and the
# server's replies reach the client. auto_start so forwarding is active before the button fires.
uart_mock:
- id: virtual_uart_server
baud_rate: 9600
auto_start: true
debug:
on_tx:
- then:
- uart_mock.inject_rx:
id: virtual_uart_client
data: !lambda return data;
- id: virtual_uart_client
baud_rate: 9600
auto_start: true
debug:
on_tx:
- then:
- uart_mock.inject_rx:
id: virtual_uart_server
data: !lambda return data;
globals:
- id: stored_1
type: uint16_t
initial_value: "0"
modbus:
- uart_id: virtual_uart_server
id: virtual_modbus_server
role: server
- uart_id: virtual_uart_client
id: virtual_modbus_client
role: client
turnaround_time: 10ms
modbus_server:
- address: 1
modbus_id: virtual_modbus_server
registers:
# Writable + readable register: the read publishes what it returns, so the test can confirm the
# write half of the 0x17 ran before the read half (Modbus 6.17).
- address: 0x01
value_type: U_WORD
read_lambda: |-
id(srv_read_1).publish_state(id(stored_1));
return id(stored_1);
write_lambda: |-
id(stored_1) = x;
id(srv_write_1).publish_state(x);
return true;
# Read-only register, returned together with 0x01 by the 2-register read half.
- address: 0x02
value_type: U_WORD
read_lambda: return 0x00AA;
sensor:
# Server-side observations.
- platform: template
name: "srv_write_1"
id: srv_write_1
- platform: template
name: "srv_read_1"
id: srv_read_1
# Client-side read-back: the values the client's on_response received.
- platform: template
name: "client_read_0"
id: client_read_0
- platform: template
name: "client_read_1"
id: client_read_1
button:
- platform: template
name: "Start Scenario"
id: start_scenario_btn
on_press:
# FC 0x17: write reg 0x0001 = 0x1234, then read regs 0x0001..0x0002 back in the same transaction.
- modbus_client.read_write_multiple_registers:
address: 0x01
read_address: 0x0001
read_count: 2
write_address: 0x0001
values: [0x1234]
on_response:
then:
- lambda: |-
// values is the read-back block: reg 0x0001 (must be the just-written 0x1234) and reg 0x0002.
if (values.size() >= 2) {
id(client_read_0).publish_state(values[0]);
id(client_read_1).publish_state(values[1]);
}
@@ -1,88 +0,0 @@
esphome:
name: uart-mock-modbus-custom-pdu
host:
api:
logger:
level: VERBOSE
external_components:
- source:
type: local
path: EXTERNAL_COMPONENT_PATH
# Dummy uart entry to satisfy modbus's DEPENDENCIES = ["uart"]
# The actual UART bus used is the uart_mock component below
uart:
baud_rate: 115200
port: /dev/null
uart_mock:
- id: virtual_uart_server
baud_rate: 9600
auto_start: true
debug:
on_tx:
- then:
- uart_mock.inject_rx:
id: virtual_uart_controller
data: !lambda return data;
- id: virtual_uart_controller
baud_rate: 9600
auto_start: true
debug:
on_tx:
- then:
- uart_mock.inject_rx:
id: virtual_uart_server
data: !lambda return data;
modbus:
- uart_id: virtual_uart_server
id: virtual_modbus_server
role: server
- uart_id: virtual_uart_controller
id: virtual_modbus_controller
role: client
turnaround_time: 10ms
modbus_controller:
- address: 1
modbus_id: virtual_modbus_controller
id: modbus_controller_1
update_interval: 1s
modbus_server:
- address: 1
modbus_id: virtual_modbus_server
id: modbus_server_1
registers:
- address: 0x01
value_type: U_WORD
read_lambda: return 259;
sensor:
# Plain read to confirm the controller <-> server link is up.
- platform: modbus_controller
modbus_controller_id: modbus_controller_1
name: "plain_read"
address: 0x01
register_type: holding
value_type: U_WORD
# Custom PDU: read holding register 0x0001, count 1. The PDU is
# {function code, address hi, address lo, count hi, count lo}; the device
# address and CRC are added by the hub. The lambda parses the response payload
# (the register value, big-endian).
- platform: modbus_controller
modbus_controller_id: modbus_controller_1
name: "custom_read"
custom_pdu: [0x03, 0x00, 0x01, 0x00, 0x01]
lambda: |-
if (data.size() < 2) return {};
return (float) ((data[0] << 8) | data[1]);
button:
- platform: template
name: "Start Scenario"
id: start_scenario_btn
# This test does not have anything to start (mock is autostart)
@@ -1,106 +0,0 @@
esphome:
name: uart-mock-modbus-dep-buffer
host:
api:
logger:
level: VERBOSE
external_components:
- source:
type: local
path: EXTERNAL_COMPONENT_PATH
# Dummy uart entry to satisfy modbus's DEPENDENCIES = ["uart"]
# The actual UART bus used is the uart_mock component below
uart:
baud_rate: 115200
port: /dev/null
uart_mock:
- id: virtual_uart_server
baud_rate: 9600
auto_start: true
debug:
on_tx:
- then:
- uart_mock.inject_rx:
id: virtual_uart_controller
data: !lambda return data;
- id: virtual_uart_controller
baud_rate: 9600
auto_start: true
debug:
on_tx:
- then:
- uart_mock.inject_rx:
id: virtual_uart_server
data: !lambda return data;
globals:
- id: reg10
type: uint16_t
initial_value: "0"
modbus:
- uart_id: virtual_uart_server
id: virtual_modbus_server
role: server
- uart_id: virtual_uart_controller
id: virtual_modbus_controller
role: client
turnaround_time: 10ms
modbus_controller:
- address: 1
modbus_id: virtual_modbus_controller
id: modbus_controller_1
update_interval: 1s
modbus_server:
- address: 1
modbus_id: virtual_modbus_server
id: modbus_server_1
registers:
- address: 0x10
value_type: U_WORD
read_lambda: return id(reg10);
write_lambda: |-
id(reg10) = x;
return true;
# A number whose write_lambda uses the DEPRECATED buffer parameter (fills `payload` with a legacy raw
# frame as words: device address + function code + data) instead of the new item->write_* API. The write
# must still land with its legacy semantics, and the one-time deprecation warning must fire only once per
# entity no matter how many writes happen.
number:
- platform: modbus_controller
modbus_controller_id: modbus_controller_1
name: "buf_number"
id: buf_number
address: 0x10
register_type: holding
value_type: U_WORD
min_value: 0
max_value: 1000
step: 1
write_lambda: |-
// Legacy raw frame as words: [addr 0x01 | fc 0x06], register 0x0010, value.
payload.push_back(0x0106);
payload.push_back(0x0010);
payload.push_back((uint16_t) x);
return {};
# Reports the server-side register so the test can observe that the deprecated buffer write landed.
sensor:
- platform: template
name: "written_value"
id: written_value
update_interval: 0.5s
lambda: "return id(reg10);"
button:
- platform: template
name: "Start Scenario"
id: start_scenario_btn
# The test drives the writes via number_command; the mock is autostart.
@@ -1,95 +0,0 @@
esphome:
name: uart-mock-modbus-lambda-invert
host:
api:
logger:
level: VERBOSE
external_components:
- source:
type: local
path: EXTERNAL_COMPONENT_PATH
# Dummy uart entry to satisfy modbus's DEPENDENCIES = ["uart"]
# The actual UART bus used is the uart_mock component below
uart:
baud_rate: 115200
port: /dev/null
uart_mock:
- id: virtual_uart_server
baud_rate: 9600
auto_start: true
debug:
on_tx:
- then:
- uart_mock.inject_rx:
id: virtual_uart_controller
data: !lambda return data;
- id: virtual_uart_controller
baud_rate: 9600
auto_start: true
debug:
on_tx:
- then:
- uart_mock.inject_rx:
id: virtual_uart_server
data: !lambda return data;
globals:
- id: reg40
type: uint16_t
initial_value: "5"
modbus:
- uart_id: virtual_uart_server
id: virtual_modbus_server
role: server
- uart_id: virtual_uart_controller
id: virtual_modbus_controller
role: client
turnaround_time: 10ms
modbus_controller:
- address: 1
modbus_id: virtual_modbus_controller
id: modbus_controller_1
update_interval: 1s
modbus_server:
- address: 1
modbus_id: virtual_modbus_server
id: modbus_server_1
registers:
- address: 0x40
value_type: U_WORD
read_lambda: return id(reg40);
write_lambda: id(reg40) = x; return true;
# An active-low holding switch: the write_lambda inverts the wire value, but the entity must still
# report the REQUESTED state. assumed_state keeps the register unpolled, so the published state comes
# only from write_state() - turning ON writes 0x0000 yet the switch shows ON.
switch:
- platform: modbus_controller
modbus_controller_id: modbus_controller_1
name: "invert_switch"
register_type: holding
address: 0x40
assumed_state: true
write_lambda: |-
return !x;
sensor:
- platform: modbus_controller
modbus_controller_id: modbus_controller_1
name: "reg_40"
address: 0x40
register_type: holding
value_type: U_WORD
button:
- platform: template
name: "Start Scenario"
id: start_scenario_btn
# This test does not have anything to start (mock is autostart)
@@ -1,97 +0,0 @@
esphome:
name: uart-mock-modbus-lambda-write
host:
api:
logger:
level: VERBOSE
external_components:
- source:
type: local
path: EXTERNAL_COMPONENT_PATH
# Dummy uart entry to satisfy modbus's DEPENDENCIES = ["uart"]
# The actual UART bus used is the uart_mock component below
uart:
baud_rate: 115200
port: /dev/null
uart_mock:
- id: virtual_uart_server
baud_rate: 9600
auto_start: true
debug:
on_tx:
- then:
- uart_mock.inject_rx:
id: virtual_uart_controller
data: !lambda return data;
- id: virtual_uart_controller
baud_rate: 9600
auto_start: true
debug:
on_tx:
- then:
- uart_mock.inject_rx:
id: virtual_uart_server
data: !lambda return data;
globals:
- id: reg30
type: uint16_t
initial_value: "0"
modbus:
- uart_id: virtual_uart_server
id: virtual_modbus_server
role: server
- uart_id: virtual_uart_controller
id: virtual_modbus_controller
role: client
turnaround_time: 10ms
modbus_controller:
- address: 1
modbus_id: virtual_modbus_controller
id: modbus_controller_1
update_interval: 1s
modbus_server:
- address: 1
modbus_id: virtual_modbus_server
id: modbus_server_1
registers:
- address: 0x30
value_type: U_WORD
read_lambda: return id(reg30);
write_lambda: id(reg30) = x; return true;
# A COIL-type switch (assumed_state, write-only) whose write_lambda ignores its own coil type and instead
# drives a HOLDING-REGISTER write on the mock server through the entity itself: `item` IS the command, so
# item->write_single_register() sends a register write from a coil entity (cross-type). Returning nothing
# (an empty optional) tells the write path the lambda already dispatched the frame - no default coil write.
switch:
- platform: modbus_controller
modbus_controller_id: modbus_controller_1
name: "cross_switch"
register_type: coil
address: 0x00
assumed_state: true
write_lambda: |-
item->write_single_register(0x30, x ? 1234 : 0);
return {};
sensor:
- platform: modbus_controller
modbus_controller_id: modbus_controller_1
name: "reg_30"
address: 0x30
register_type: holding
value_type: U_WORD
button:
- platform: template
name: "Start Scenario"
id: start_scenario_btn
# This test does not have anything to start (mock is autostart)
@@ -0,0 +1,233 @@
esphome:
name: uart-mock-modbus-loopback
host:
api:
logger:
level: VERBOSE
external_components:
- source:
type: local
path: EXTERNAL_COMPONENT_PATH
# Dummy uart entry to satisfy modbus's DEPENDENCIES = ["uart"]
# The actual UART bus used is the uart_mock component below
uart:
baud_rate: 115200
port: /dev/null
# Shared loopback fixture (see the shared_yaml markers in the test file);
# register spaces are disjoint so each test only observes its own entities.
uart_mock:
- id: virtual_uart_server
baud_rate: 9600
auto_start: true
debug:
on_tx:
- then:
- uart_mock.inject_rx:
id: virtual_uart_controller
data: !lambda return data;
- id: virtual_uart_controller
baud_rate: 9600
auto_start: true
debug:
on_tx:
- then:
- uart_mock.inject_rx:
id: virtual_uart_server
data: !lambda return data;
globals:
- id: reg10
type: uint16_t
initial_value: "100"
- id: reg11
type: uint16_t
initial_value: "200"
- id: reg12
type: uint16_t
initial_value: "300"
- id: reg13
type: uint16_t
initial_value: "0xABCD"
- id: reg30
type: uint16_t
initial_value: "0"
- id: reg40
type: uint16_t
initial_value: "5"
- id: reg50
type: uint16_t
initial_value: "0"
modbus:
- uart_id: virtual_uart_server
id: virtual_modbus_server
role: server
- uart_id: virtual_uart_controller
id: virtual_modbus_controller
role: client
turnaround_time: 10ms
modbus_controller:
- address: 1
modbus_id: virtual_modbus_controller
id: modbus_controller_1
update_interval: 1s
modbus_server:
- address: 1
modbus_id: virtual_modbus_server
registers:
- address: 0x01
value_type: U_WORD
read_lambda: return 259;
- address: 0x10
value_type: U_WORD
read_lambda: return id(reg10);
write_lambda: id(reg10) = x; return true;
- address: 0x11
value_type: U_WORD
read_lambda: return id(reg11);
write_lambda: id(reg11) = x; return true;
- address: 0x12
value_type: U_WORD
read_lambda: return id(reg12);
write_lambda: id(reg12) = x; return true;
- address: 0x13
value_type: U_WORD
read_lambda: return id(reg13);
- address: 0x30
value_type: U_WORD
read_lambda: return id(reg30);
write_lambda: id(reg30) = x; return true;
- address: 0x40
value_type: U_WORD
read_lambda: return id(reg40);
write_lambda: id(reg40) = x; return true;
- address: 0x50
value_type: U_WORD
read_lambda: return id(reg50);
write_lambda: id(reg50) = x; return true;
# Byte-based offset: 2 bytes -> register 0x11 (the old code folded it in as a
# register count, hitting 0x12). assumed_state keeps the switch write-only.
switch:
- platform: modbus_controller
modbus_controller_id: modbus_controller_1
name: "offset_switch"
register_type: holding
address: 0x10
offset: 2
assumed_state: true
# Reading switch, byte offset 6 -> register 0x13; the pre-fix resolution (0x16)
# would draw ILLEGAL_DATA_ADDRESS and never publish.
- platform: modbus_controller
modbus_controller_id: modbus_controller_1
name: "read_offset_switch"
register_type: holding
address: 0x10
offset: 6
bitmask: 0x1
# Coil switch whose write_lambda dispatches a holding-register write via `item`;
# returning an empty optional suppresses the default coil write.
- platform: modbus_controller
modbus_controller_id: modbus_controller_1
name: "cross_switch"
register_type: coil
address: 0x00
assumed_state: true
write_lambda: |-
item->write_single_register(0x30, x ? 1234 : 0);
return {};
# Active-low: the write_lambda inverts the wire value but the entity must still
# report the requested state (assumed_state keeps the register unpolled).
- platform: modbus_controller
modbus_controller_id: modbus_controller_1
name: "invert_switch"
register_type: holding
address: 0x40
assumed_state: true
write_lambda: |-
return !x;
# Uses the deprecated buffer parameter (legacy raw frame as words); the write
# must land and the deprecation warning must fire only once per entity.
number:
- platform: modbus_controller
modbus_controller_id: modbus_controller_1
name: "buf_number"
id: buf_number
address: 0x50
register_type: holding
value_type: U_WORD
min_value: 0
max_value: 1000
step: 1
write_lambda: |-
// Legacy raw frame as words: [addr 0x01 | fc 0x06], register 0x0050, value.
payload.push_back(0x0106);
payload.push_back(0x0050);
payload.push_back((uint16_t) x);
return {};
sensor:
- platform: modbus_controller
modbus_controller_id: modbus_controller_1
name: "plain_read"
address: 0x01
register_type: holding
value_type: U_WORD
# Custom PDU: read holding register 0x0001; device address and CRC are added
# by the hub. The lambda parses the big-endian register value.
- platform: modbus_controller
modbus_controller_id: modbus_controller_1
name: "custom_read"
custom_pdu: [0x03, 0x00, 0x01, 0x00, 0x01]
lambda: |-
if (data.size() < 2) return {};
return (float) ((data[0] << 8) | data[1]);
- platform: modbus_controller
modbus_controller_id: modbus_controller_1
name: "reg_10"
address: 0x10
register_type: holding
value_type: U_WORD
- platform: modbus_controller
modbus_controller_id: modbus_controller_1
name: "reg_11"
address: 0x11
register_type: holding
value_type: U_WORD
- platform: modbus_controller
modbus_controller_id: modbus_controller_1
name: "reg_12"
address: 0x12
register_type: holding
value_type: U_WORD
- platform: modbus_controller
modbus_controller_id: modbus_controller_1
name: "reg_30"
address: 0x30
register_type: holding
value_type: U_WORD
- platform: modbus_controller
modbus_controller_id: modbus_controller_1
name: "reg_40"
address: 0x40
register_type: holding
value_type: U_WORD
# Reports the server-side register so the test can observe that the deprecated buffer write landed.
- platform: template
name: "written_value"
id: written_value
update_interval: 0.5s
lambda: "return id(reg50);"
button:
- platform: template
name: "Start Scenario"
id: start_scenario_btn
# Nothing to start (mock is autostart); tests drive entities directly
@@ -1,5 +1,5 @@
esphome:
name: uart-mock-modbus-server-contro
name: uart-mock-modbus-mesh
host:
api:
@@ -17,13 +17,14 @@ uart:
baud_rate: 115200
port: /dev/null
# Shared 3-bus mesh (see the shared_yaml markers): addr 1 = typed read-only
# registers, addr 5 = the read/write 0x17 target, addr 2/3 on the second
# server hub. auto_start everywhere: the controller polls at boot, so the
# forwarding must already be live or early requests generate warnings.
# Every test presses Start Scenario, so all merged actions fire in every test.
uart_mock:
- id: virtual_uart_server
baud_rate: 9600
# auto_start must be true for loopback fixtures: the modbus controller
# polls on its update_interval immediately at boot, so the uart_mock
# forwarding must already be active or early requests are lost and
# generate modbus warnings.
auto_start: true
debug:
on_tx:
@@ -31,35 +32,68 @@ uart_mock:
- uart_mock.inject_rx:
id: virtual_uart_controller
data: !lambda return data;
- id: virtual_uart_controller
- uart_mock.inject_rx:
id: virtual_uart_server_2
data: !lambda return data;
- id: virtual_uart_server_2
baud_rate: 9600
auto_start: true # See comment on virtual_uart_server above
auto_start: true
debug:
on_tx:
- then:
- uart_mock.inject_rx:
id: virtual_uart_server
data: !lambda return data;
- uart_mock.inject_rx:
id: virtual_uart_controller
data: !lambda return data;
- id: virtual_uart_controller
baud_rate: 9600
auto_start: true
debug:
on_tx:
- then:
- uart_mock.inject_rx:
id: virtual_uart_server
data: !lambda return data;
- uart_mock.inject_rx:
id: virtual_uart_server_2
data: !lambda return data;
globals:
- id: stored_1
type: uint16_t
initial_value: "0"
modbus:
- uart_id: virtual_uart_server
id: virtual_modbus_server
role: server
- uart_id: virtual_uart_server_2
id: virtual_modbus_server_2
role: server
- uart_id: virtual_uart_controller
id: virtual_modbus_controller
id: virtual_modbus_client
role: client
turnaround_time: 10ms
modbus_controller:
- address: 1
modbus_id: virtual_modbus_controller
modbus_id: virtual_modbus_client
id: modbus_controller_1
update_interval: 1s
- address: 2
modbus_id: virtual_modbus_client
id: modbus_controller_2
update_interval: 1s
- address: 3
modbus_id: virtual_modbus_client
id: modbus_controller_3
update_interval: 1s
modbus_server:
- address: 1
modbus_id: virtual_modbus_server
id: modbus_server_1
registers:
- address: 0x01
value_type: U_WORD
@@ -103,6 +137,34 @@ modbus_server:
- address: 0x28
value_type: FP32_R
read_lambda: return 3.14;
- address: 5
modbus_id: virtual_modbus_server
registers:
# Writable + readable register: srv_write_1 plus the client's read-back
# confirm the write half of the 0x17 ran before the read half (Modbus 6.17).
- address: 0x01
value_type: U_WORD
read_lambda: return id(stored_1);
write_lambda: |-
id(stored_1) = x;
id(srv_write_1).publish_state(x);
return true;
# Read-only register, returned together with 0x01 by the 2-register read half.
- address: 0x02
value_type: U_WORD
read_lambda: return 0x00AA;
- address: 2
modbus_id: virtual_modbus_server_2
registers:
- address: 0x01
value_type: U_WORD
read_lambda: return 919;
- address: 3
modbus_id: virtual_modbus_server_2
registers:
- address: 0x01
value_type: U_WORD
read_lambda: return 929;
sensor:
- platform: modbus_controller
@@ -195,9 +257,46 @@ sensor:
address: 0x28
register_type: holding
value_type: FP32_R
- platform: modbus_controller
modbus_controller_id: modbus_controller_2
name: "multi_reg_a"
address: 0x01
register_type: holding
value_type: U_WORD
- platform: modbus_controller
modbus_controller_id: modbus_controller_3
name: "multi_reg_b"
address: 0x01
register_type: holding
value_type: U_WORD
# client_read_write observations, server- and client-side.
- platform: template
name: "srv_write_1"
id: srv_write_1
- platform: template
name: "client_read_0"
id: client_read_0
- platform: template
name: "client_read_1"
id: client_read_1
button:
- platform: template
name: "Start Scenario"
id: start_scenario_btn
# This test does not have anything to start (mock is autostart)
on_press:
# FC 0x17: write reg 0x0001 = 0x1234, then read regs 0x0001..0x0002 back in the same transaction.
- modbus_client.read_write_multiple_registers:
address: 5
read_address: 0x0001
read_count: 2
write_address: 0x0001
values: [0x1234]
on_response:
then:
- lambda: |-
// values is the read-back block: reg 0x0001 (must be the just-written 0x1234) and reg 0x0002.
if (values.size() >= 2) {
id(client_read_0).publish_state(values[0]);
id(client_read_1).publish_state(values[1]);
}
@@ -1,138 +0,0 @@
esphome:
name: uart-mock-modbus-reg-offset
host:
api:
logger:
level: VERBOSE
external_components:
- source:
type: local
path: EXTERNAL_COMPONENT_PATH
# Dummy uart entry to satisfy modbus's DEPENDENCIES = ["uart"]
# The actual UART bus used is the uart_mock component below
uart:
baud_rate: 115200
port: /dev/null
uart_mock:
- id: virtual_uart_server
baud_rate: 9600
auto_start: true
debug:
on_tx:
- then:
- uart_mock.inject_rx:
id: virtual_uart_controller
data: !lambda return data;
- id: virtual_uart_controller
baud_rate: 9600
auto_start: true
debug:
on_tx:
- then:
- uart_mock.inject_rx:
id: virtual_uart_server
data: !lambda return data;
globals:
- id: reg10
type: uint16_t
initial_value: "100"
- id: reg11
type: uint16_t
initial_value: "200"
- id: reg12
type: uint16_t
initial_value: "300"
- id: reg13
type: uint16_t
initial_value: "0xABCD"
modbus:
- uart_id: virtual_uart_server
id: virtual_modbus_server
role: server
- uart_id: virtual_uart_controller
id: virtual_modbus_controller
role: client
turnaround_time: 10ms
modbus_controller:
- address: 1
modbus_id: virtual_modbus_controller
id: modbus_controller_1
update_interval: 1s
modbus_server:
- address: 1
modbus_id: virtual_modbus_server
id: modbus_server_1
registers:
- address: 0x10
value_type: U_WORD
read_lambda: return id(reg10);
write_lambda: id(reg10) = x; return true;
- address: 0x11
value_type: U_WORD
read_lambda: return id(reg11);
write_lambda: id(reg11) = x; return true;
- address: 0x12
value_type: U_WORD
read_lambda: return id(reg12);
write_lambda: id(reg12) = x; return true;
- address: 0x13
value_type: U_WORD
read_lambda: return id(reg13);
write_lambda: id(reg13) = x; return true;
# A holding-register switch at 0x10 with a 2-BYTE offset. offset is byte-based, so the write must target
# register 0x10 + 2/2 = 0x11. The old (pre-fix) behavior folded offset into the address as a register
# count, hitting 0x12 instead. assumed_state keeps the switch write-only so it does not read any register.
switch:
- platform: modbus_controller
modbus_controller_id: modbus_controller_1
name: "offset_switch"
register_type: holding
address: 0x10
offset: 2
assumed_state: true
# A holding-register switch that READS its state. Byte offset 6 -> register 0x10 + 6/2 = 0x13. Post-fix
# the switch itself resolves to 0x13 (the even byte offset folds into the address as whole registers) and
# joins the 0x10..0x13 range, so no separate 0x13 sensor is needed. Pre-fix the whole byte offset folds
# into the address (0x16), where the server answers ILLEGAL_DATA_ADDRESS and the switch never publishes.
- platform: modbus_controller
modbus_controller_id: modbus_controller_1
name: "read_offset_switch"
register_type: holding
address: 0x10
offset: 6
bitmask: 0x1
sensor:
- platform: modbus_controller
modbus_controller_id: modbus_controller_1
name: "reg_10"
address: 0x10
register_type: holding
value_type: U_WORD
- platform: modbus_controller
modbus_controller_id: modbus_controller_1
name: "reg_11"
address: 0x11
register_type: holding
value_type: U_WORD
- platform: modbus_controller
modbus_controller_id: modbus_controller_1
name: "reg_12"
address: 0x12
register_type: holding
value_type: U_WORD
button:
- platform: template
name: "Start Scenario"
id: start_scenario_btn
# This test does not have anything to start (mock is autostart)
@@ -1,124 +0,0 @@
esphome:
name: uart-mock-modbus-server-test
host:
api:
logger:
level: VERBOSE
external_components:
- source:
type: local
path: EXTERNAL_COMPONENT_PATH
# Dummy uart entry to satisfy modbus's DEPENDENCIES = ["uart"]
# The actual UART bus used is the uart_mock component below
uart:
baud_rate: 115200
port: /dev/null
uart_mock:
- id: virtual_uart_dev
baud_rate: 9600
rx_full_threshold: 120
rx_timeout: 2
auto_start: false
debug:
injections:
- delay: 100ms
inject_rx: [0x01, 0x03, 0x00, 0x03, 0x00, 0x01, 0x74, 0x0A] # Read holding register 3 on device 1 (basic_read)
- delay: 100ms
# Read holding register 7 on device 2
# Reply from device 2
# Read holding register 5 on device 1 (read_after_peer_response)
inject_rx:
[
0x02,
0x03,
0x00,
0x07,
0x00,
0x01,
0x35,
0xF8,
0x02,
0x03,
0x02,
0x00,
0xF0,
0xFC,
0x00,
0x01,
0x03,
0x00,
0x05,
0x00,
0x01,
0x94,
0x0B,
]
- delay: 100ms
inject_rx: [0x02, 0x03, 0x00, 0x07, 0x00, 0x01, 0x35, 0xF8] # Read holding register 7 on device 2, with no response
- delay: 100ms
# Read holding register 7 on device 2, with no response
# Read holding register A on device 1 (read_after_peer_timeout)
inject_rx:
[
0x02,
0x03,
0x00,
0x07,
0x00,
0x01,
0x35,
0xF8,
0x01,
0x03,
0x00,
0x0A,
0x00,
0x01,
0xA4,
0x08,
]
modbus:
uart_id: virtual_uart_dev
role: server
modbus_server:
- address: 1
registers:
- address: 0x03
value_type: U_WORD
read_lambda: |-
id(basic_read).publish_state(1);
return 1;
- address: 0x05
value_type: U_WORD
read_lambda: |-
id(read_after_peer_response).publish_state(1);
return 1;
- address: 0x0A
value_type: U_WORD
read_lambda: |-
id(read_after_peer_timeout).publish_state(1);
return 1;
sensor:
- platform: template
name: "basic_read"
id: basic_read
- platform: template
name: "read_after_peer_response"
id: read_after_peer_response
- platform: template
name: "read_after_peer_timeout"
id: read_after_peer_timeout
button:
- platform: template
name: "Start Scenario"
id: start_scenario_btn
on_press:
- lambda: "id(virtual_uart_dev).start_scenario();"
@@ -1,116 +0,0 @@
esphome:
name: uart-mock-modbus-server-mult
host:
api:
logger:
level: VERBOSE
external_components:
- source:
type: local
path: EXTERNAL_COMPONENT_PATH
# Dummy uart entry to satisfy modbus's DEPENDENCIES = ["uart"]
# The actual UART bus used is the uart_mock component below
uart:
baud_rate: 115200
port: /dev/null
uart_mock:
- id: virtual_uart_server
baud_rate: 9600
# auto_start must be true for loopback fixtures: the modbus controller
# polls on its update_interval immediately at boot, so the uart_mock
# forwarding must already be active or early requests are lost and
# generate modbus warnings.
auto_start: true
debug:
on_tx:
- then:
- uart_mock.inject_rx:
id: virtual_uart_controller
data: !lambda return data;
- uart_mock.inject_rx:
id: virtual_uart_server_2
data: !lambda return data;
- id: virtual_uart_server_2
baud_rate: 9600
auto_start: true # See comment on virtual_uart_server above
debug:
on_tx:
- then:
- uart_mock.inject_rx:
id: virtual_uart_server
data: !lambda return data;
- uart_mock.inject_rx:
id: virtual_uart_controller
data: !lambda return data;
- id: virtual_uart_controller
baud_rate: 9600
auto_start: true # See comment on virtual_uart_server above
debug:
on_tx:
- then:
- uart_mock.inject_rx:
id: virtual_uart_server
data: !lambda return data;
- uart_mock.inject_rx:
id: virtual_uart_server_2
data: !lambda return data;
modbus:
- uart_id: virtual_uart_server
id: virtual_modbus_server
role: server
- uart_id: virtual_uart_server_2
id: virtual_modbus_server_2
role: server
- uart_id: virtual_uart_controller
id: virtual_modbus_client
role: client
turnaround_time: 10ms
modbus_controller:
- address: 1
modbus_id: virtual_modbus_client
update_interval: 1s
id: modbus_controller_1
- address: 2
modbus_id: virtual_modbus_client
update_interval: 1s
id: modbus_controller_2
modbus_server:
- address: 1
modbus_id: virtual_modbus_server
registers:
- address: 0x01
value_type: U_WORD
read_lambda: return 919;
- address: 2
modbus_id: virtual_modbus_server_2
registers:
- address: 0x01
value_type: U_WORD
read_lambda: return 929;
sensor:
- platform: modbus_controller
modbus_controller_id: modbus_controller_1
name: "reg_u_word"
address: 0x01
register_type: holding
value_type: U_WORD
- platform: modbus_controller
modbus_controller_id: modbus_controller_2
name: "reg_u_word_2"
address: 0x01
register_type: holding
value_type: U_WORD
button:
- platform: template
name: "Start Scenario"
id: start_scenario_btn
# This test does not have anything to start (mock is autostart)
@@ -1,5 +1,5 @@
esphome:
name: uart-mock-modbus-srv-rw
name: uart-mock-modbus-srv-injected
host:
api:
@@ -17,6 +17,8 @@ uart:
baud_rate: 115200
port: /dev/null
# Shared server-role fixture (see the shared_yaml markers in the test file);
# the injections concatenate and each test waits only on its own sensors.
uart_mock:
- id: virtual_uart_dev
baud_rate: 9600
@@ -25,18 +27,31 @@ uart_mock:
auto_start: false
debug:
injections:
# FC 0x17 Read/Write Multiple Registers on device 1:
# write reg 0x0001 = 0x1234 (qty 1), then read regs 0x0001..0x0002 (qty 2).
# Per Modbus 6.17 the write is performed before the read, so reg 0x0001 must
# read back the just-written 0x1234 in the same request.
- delay: 100ms
inject_rx: [0x01, 0x03, 0x00, 0x03, 0x00, 0x01, 0x74, 0x0A] # Read holding register 3 on device 1 (basic_read)
- delay: 100ms
# Read holding register 7 on device 2, its reply, then read holding
# register 5 on device 1 (read_after_peer_response)
inject_rx: [0x02, 0x03, 0x00, 0x07, 0x00, 0x01, 0x35, 0xF8,
0x02, 0x03, 0x02, 0x00, 0xF0, 0xFC,
0x00, 0x01, 0x03, 0x00, 0x05, 0x00, 0x01, 0x94, 0x0B]
- delay: 100ms
inject_rx: [0x02, 0x03, 0x00, 0x07, 0x00, 0x01, 0x35, 0xF8] # Read holding register 7 on device 2, with no response
- delay: 100ms
# Read holding register 7 on device 2 with no response, then read
# holding register A on device 1 (read_after_peer_timeout)
inject_rx: [0x02, 0x03, 0x00, 0x07, 0x00, 0x01, 0x35, 0xF8,
0x01, 0x03, 0x00, 0x0A, 0x00, 0x01, 0xA4, 0x08]
# FC 0x17 on device 1: write reg 0x0001 = 0x1234 then read 0x0001..0x0002;
# per Modbus 6.17 the write runs first, so 0x0001 must read back 0x1234.
- delay: 100ms
inject_rx:
[0x01, 0x17, 0x00, 0x01, 0x00, 0x02, 0x00, 0x01, 0x00, 0x01, 0x02, 0x12, 0x34, 0x49, 0xD8]
# FC 0x17: write reg 0x0003 = 0x5678 (qty 1), then read reg 0x0003 (qty 1) -
# FC 0x17: write reg 0x0006 = 0x5678 (qty 1), then read reg 0x0006 (qty 1) -
# a write and read targeting a different register block.
- delay: 100ms
inject_rx:
[0x01, 0x17, 0x00, 0x03, 0x00, 0x01, 0x00, 0x03, 0x00, 0x01, 0x02, 0x56, 0x78, 0x9B, 0x10]
[0x01, 0x17, 0x00, 0x06, 0x00, 0x01, 0x00, 0x06, 0x00, 0x01, 0x02, 0x56, 0x78, 0x8B, 0x55]
globals:
- id: stored_1
@@ -70,8 +85,18 @@ modbus_server:
read_lambda: |-
id(rw_read_2).publish_state(0x00AA);
return 0x00AA;
# Second writable + readable register, targeted by the second request.
- address: 0x03
value_type: U_WORD
read_lambda: |-
id(basic_read).publish_state(1);
return 1;
- address: 0x05
value_type: U_WORD
read_lambda: |-
id(read_after_peer_response).publish_state(1);
return 1;
# Second writable + readable register, targeted by the second FC 0x17 request.
- address: 0x06
value_type: U_WORD
read_lambda: |-
id(rw_read_3).publish_state(id(stored_3));
@@ -80,8 +105,22 @@ modbus_server:
id(stored_3) = x;
id(rw_write_3).publish_state(x);
return true;
- address: 0x0A
value_type: U_WORD
read_lambda: |-
id(read_after_peer_timeout).publish_state(1);
return 1;
sensor:
- platform: template
name: "basic_read"
id: basic_read
- platform: template
name: "read_after_peer_response"
id: read_after_peer_response
- platform: template
name: "read_after_peer_timeout"
id: read_after_peer_timeout
- platform: template
name: "rw_write_1"
id: rw_write_1
+11 -3
View File
@@ -1,7 +1,7 @@
"""Helpers for manipulating the host platform's preferences file.
ESPHome's host platform stores preferences in
``~/.esphome/prefs/<app_name>.prefs`` using a simple binary layout that
``$ESPHOME_PREFDIR/<app_name>.prefs`` using a simple binary layout that
mirrors ``HostPreferences::sync()``:
``[uint32_t key][uint8_t len][uint8_t data[len]]`` per entry.
@@ -11,13 +11,21 @@ boot (e.g. forcing safe mode) or to clear stale state between runs.
from __future__ import annotations
import os
from pathlib import Path
import struct
def host_prefs_path(device_name: str) -> Path:
"""Return the on-disk prefs file path for a host-platform device."""
return Path.home() / ".esphome" / "prefs" / f"{device_name}.prefs"
"""Return the on-disk prefs file path for a host-platform device.
Requires ESPHOME_PREFDIR, which the autouse isolated_preferences fixture
sets; refusing the ~/.esphome/prefs fallback keeps tests off real user
data if the fixture is ever bypassed."""
prefdir = os.environ.get("ESPHOME_PREFDIR")
if not prefdir:
raise RuntimeError("ESPHOME_PREFDIR is not set; refusing the real prefs dir")
return Path(prefdir) / f"{device_name}.prefs"
def clear_host_prefs(device_name: str) -> None:
+4 -5
View File
@@ -125,12 +125,11 @@ class RawApiClient:
await self.read_until_frame(MESSAGE_TYPE_OF[api_pb2.HelloResponse])
async def send_message(self, msg: message.Message) -> None:
await self.send_raw(MESSAGE_TYPE_OF[type(msg)], msg.SerializeToString())
async def send_raw(self, msg_type: int, payload: bytes) -> None:
"""Send a frame with a hand built payload, for shapes protobuf will not serialize."""
loop = asyncio.get_running_loop()
await loop.sock_sendall(self._sock, encode_frame(msg_type, payload))
await loop.sock_sendall(
self._sock,
encode_frame(MESSAGE_TYPE_OF[type(msg)], msg.SerializeToString()),
)
async def read_until_frame(self, msg_type: int, timeout: float = 10.0) -> None:
"""Read until at least one frame of msg_type has been received."""
-40
View File
@@ -57,46 +57,6 @@ async def wait_for_state(
return await asyncio.wait_for(future, timeout=timeout)
class StateWaiter:
"""Route one state subscription to any number of predicate waits."""
def __init__(self) -> None:
self._waiters: list[
tuple[Callable[[EntityState], bool], asyncio.Future[EntityState]]
] = []
def on_state(self, state: EntityState) -> None:
for predicate, future in self._waiters:
if future.done():
continue
try:
matched = predicate(state)
except Exception as exc: # noqa: BLE001 the wait re-raises it, the callback must not die
future.set_exception(exc)
continue
if matched:
future.set_result(state)
async def expect(
self,
predicate: Callable[[EntityState], bool],
timeout: float = 5.0,
label: str | None = None,
) -> EntityState:
"""Wait for the next state matching ``predicate``; states seen before this call do not count."""
entry = (predicate, asyncio.get_running_loop().create_future())
self._waiters.append(entry)
try:
async with asyncio.timeout(timeout):
return await entry[1]
except TimeoutError:
raise TimeoutError(
f"no state matched {label or predicate} within {timeout}s"
) from None
finally:
self._waiters.remove(entry)
def find_entity[T: EntityInfo](
entities: list[EntityInfo],
object_id_substring: str,
@@ -1,142 +0,0 @@
"""decode_field() must take fields that match their declared wire type, drop the ones that do
not, skip unknown fields, and handle two byte tags, varints and length prefixes."""
from __future__ import annotations
from collections.abc import Callable
import struct
from aioesphomeapi import (
EntityState,
LightState,
NumberState,
SwitchState,
TextState,
api_pb2,
)
import pytest
from .raw_api_client import MESSAGE_TYPE_OF, RawApiClient, encode_varint
from .state_utils import InitialStateHelper, StateWaiter, require_entity
from .types import APIClientConnectedFactory, RunCompiledFunction
SWITCH_COMMAND = MESSAGE_TYPE_OF[api_pb2.SwitchCommandRequest]
WIRE_VARINT, WIRE_LENGTH, WIRE_FIXED32 = 0, 2, 5
def tag(field: int, wire_type: int) -> bytes:
return encode_varint((field << 3) | wire_type)
@pytest.mark.asyncio
async def test_api_decode_wire_types(
yaml_config: str,
run_compiled: RunCompiledFunction,
api_client_connected: APIClientConnectedFactory,
unused_tcp_port: int,
) -> None:
async with (
run_compiled(yaml_config),
api_client_connected() as client,
RawApiClient(unused_tcp_port) as raw,
):
entities, _ = await client.list_entities_services()
switch = require_entity(entities, "wire_switch")
light = require_entity(entities, "wire_light")
text = require_entity(entities, "wire_text")
number = require_entity(entities, "wire_number")
key = tag(1, WIRE_FIXED32) + struct.pack("<I", switch.key)
on, off = tag(2, WIRE_VARINT) + b"\x01", tag(2, WIRE_VARINT) + b"\x00"
switch_states: list[bool] = []
waiter = StateWaiter()
def on_state(state: EntityState) -> None:
if isinstance(state, SwitchState) and state.key == switch.key:
switch_states.append(state.state)
waiter.on_state(state)
def switch_is(value: bool) -> Callable[[EntityState], bool]:
return lambda s: (
isinstance(s, SwitchState) and s.key == switch.key and s.state is value
)
def number_is(value: float) -> Callable[[EntityState], bool]:
return lambda s: (
isinstance(s, NumberState) and s.key == number.key and s.state == value
)
initial = InitialStateHelper(entities)
client.subscribe_states(initial.on_state_wrapper(on_state))
await initial.wait_for_initial_states()
await raw.connect()
# A well formed command: fixed32 key, varint state
await raw.send_raw(SWITCH_COMMAND, key + on)
await waiter.expect(switch_is(True))
await raw.send_raw(SWITCH_COMMAND, key + off)
await waiter.expect(switch_is(False))
# The same field with the wrong wire type is dropped, and a varint key never matches an
# entity; each of these would turn the switch on if the payload were read as a varint
seen = len(switch_states)
await raw.send_raw(SWITCH_COMMAND, key + tag(2, WIRE_LENGTH) + b"\x01\x01")
await raw.send_raw(
SWITCH_COMMAND, key + tag(2, WIRE_FIXED32) + b"\x01\x00\x00\x00"
)
await raw.send_raw(
SWITCH_COMMAND, tag(1, WIRE_VARINT) + encode_varint(switch.key) + on
)
# Ordered on the raw socket itself: this frame cannot be parsed before the bad ones, so
# the only switch state since the marker must be the one it produces
await raw.send_raw(SWITCH_COMMAND, key + on)
await waiter.expect(switch_is(True), label="switch on after wrong wire types")
assert switch_states[seen:] == [True]
await raw.send_raw(SWITCH_COMMAND, key + off)
await waiter.expect(switch_is(False))
# Truncated bodies stop the decode loop without taking the connection down: a tag with its
# continuation bit set and nothing after it, a length prefix past the end of the payload,
# and a fixed32 with two of its four bytes
seen = len(switch_states)
await raw.send_raw(SWITCH_COMMAND, key + b"\x80")
await raw.send_raw(SWITCH_COMMAND, key + tag(2, WIRE_LENGTH) + b"\x7f" + b"ab")
await raw.send_raw(SWITCH_COMMAND, tag(1, WIRE_FIXED32) + b"\x01\x02")
await raw.send_raw(SWITCH_COMMAND, key + on)
await waiter.expect(switch_is(True), label="switch on after truncated frames")
assert switch_states[seen:] == [True]
await raw.send_raw(SWITCH_COMMAND, key + off)
await waiter.expect(switch_is(False))
# A negative number goes through the fixed32 float path of a normal client
client.number_command(number.key, -77.5)
await waiter.expect(number_is(-77.5))
# An unknown field ahead of the known ones is skipped; field 200 needs a two byte tag
await raw.send_raw(
SWITCH_COMMAND, tag(200, WIRE_VARINT) + encode_varint(300) + key + on
)
await waiter.expect(switch_is(True))
# Two byte tags (effect fields 18 and 19) and a two byte varint (300 ms transition)
client.light_command(
light.key, state=True, brightness=0.5, transition_length=0.3, effect="Pulse"
)
await waiter.expect(
lambda s: (
isinstance(s, LightState) and s.key == light.key and s.effect == "Pulse"
)
)
client.light_command(light.key, effect="None", state=False)
await waiter.expect(
lambda s: isinstance(s, LightState) and s.key == light.key and not s.state
)
# A string whose length prefix needs two varint bytes
long_text = "w" * 200
client.text_command(text.key, long_text)
await waiter.expect(
lambda s: (
isinstance(s, TextState) and s.key == text.key and s.state == long_text
)
)
@@ -1,37 +0,0 @@
"""Messages without fields go through the shared ProtoMessage entry points on both directions."""
from __future__ import annotations
from aioesphomeapi import api_pb2
import pytest
from .raw_api_client import MESSAGE_TYPE_OF, RawApiClient
from .types import RunCompiledFunction
@pytest.mark.asyncio
async def test_api_empty_message_roundtrip(
yaml_config: str,
run_compiled: RunCompiledFunction,
unused_tcp_port: int,
) -> None:
async with run_compiled(yaml_config), RawApiClient(unused_tcp_port) as client:
await client.connect()
# Field free request and reply on the plain send path
await client.send_message(api_pb2.PingRequest())
await client.read_until_frame(MESSAGE_TYPE_OF[api_pb2.PingResponse])
# Field free request answered by a message with fields, and a list that ends with
# the field free ListEntitiesDoneResponse through the batching path
await client.send_message(api_pb2.DeviceInfoRequest())
await client.read_until_frame(MESSAGE_TYPE_OF[api_pb2.DeviceInfoResponse])
await client.send_message(api_pb2.ListEntitiesRequest())
await client.read_until_frame(MESSAGE_TYPE_OF[api_pb2.ListEntitiesDoneResponse])
assert (
client.frame_counts[MESSAGE_TYPE_OF[api_pb2.ListEntitiesSwitchResponse]]
== 1
)
await client.send_message(api_pb2.DisconnectRequest())
await client.read_until_frame(MESSAGE_TYPE_OF[api_pb2.DisconnectResponse])
@@ -1,78 +0,0 @@
"""Encode paths at their branch boundaries: zero skipped float, fixed32 state, negative int32,
length prefixes of two varint bytes and two byte field tags."""
from __future__ import annotations
import asyncio
from aioesphomeapi import (
NumberState,
SelectInfo,
SensorInfo,
SensorState,
TextSensorState,
)
import pytest
from .state_utils import InitialStateHelper, StateWaiter, require_entity
from .types import APIClientConnectedFactory, RunCompiledFunction
LONG_OPTION = (
"option-with-a-name-long-enough-that-its-length-prefix-needs-two-varint-bytes-"
"when-the-list-entities-response-is-encoded-xxxxxxxxxx"
)
@pytest.mark.asyncio
async def test_api_encode_boundaries(
yaml_config: str,
run_compiled: RunCompiledFunction,
api_client_connected: APIClientConnectedFactory,
) -> None:
async with run_compiled(yaml_config), api_client_connected() as client:
device_info, (entities, _) = await asyncio.gather(
client.device_info(), client.list_entities_services()
)
assert device_info.suggested_area == "Kitchen"
sensor = require_entity(entities, "zero_then_value", SensorInfo)
assert sensor.accuracy_decimals == -2
select = require_entity(entities, "long_option_select", SelectInfo)
assert len(LONG_OPTION) >= 128
assert select.options == ["short", LONG_OPTION]
text = require_entity(entities, "long_text")
number = require_entity(entities, "negative_number")
button = require_entity(entities, "publish_values")
initial = InitialStateHelper(entities)
waiter = StateWaiter()
client.subscribe_states(initial.on_state_wrapper(waiter.on_state))
await initial.wait_for_initial_states()
# A float of exactly zero is skipped on the wire and must still read as 0.0, not missing
first = initial.initial_states[sensor.key]
assert isinstance(first, SensorState)
assert first.state == 0.0 and not first.missing_state
first_number = initial.initial_states[number.key]
assert isinstance(first_number, NumberState)
assert first_number.state == -123.5
client.button_command(button.key)
await asyncio.gather(
waiter.expect(
lambda s: (
isinstance(s, SensorState)
and s.key == sensor.key
and s.state == 12.5
),
label="sensor 12.5",
),
waiter.expect(
lambda s: (
isinstance(s, TextSensorState)
and s.key == text.key
and s.state == "y" * 200
),
label="text 200 x y",
),
)
@@ -0,0 +1,58 @@
"""Integration test for noise session resume."""
from __future__ import annotations
import asyncio
import aioesphomeapi.core
import pytest
from .types import APIClientConnectedFactory, RunCompiledFunction
NOISE_KEY = "N4Yle5YirwZhPiHHsdZLdOA73ndj/84veVaLhTvxCuU="
@pytest.mark.asyncio
async def test_api_noise_resume(
yaml_config: str,
run_compiled: RunCompiledFunction,
api_client_connected: APIClientConnectedFactory,
) -> None:
"""A reconnect with the ticket from the first connection resumes the session."""
if not hasattr(aioesphomeapi.core, "ResumeAPIError"):
pytest.skip("aioesphomeapi without noise session resume")
resumed = asyncio.Event()
resumed_count = 0
def on_line(line: str) -> None:
nonlocal resumed_count
if "Session resumed" in line:
resumed_count += 1
resumed.set()
async with (
run_compiled(yaml_config, line_callback=on_line),
api_client_connected(noise_psk=NOISE_KEY) as client,
):
# First connection: full handshake, the device issues a ticket
info = await client.device_info()
assert info.name == "host-noise-resume"
assert resumed_count == 0
# Same client reconnects and offers the ticket
await client.disconnect()
await client.connect(login=True)
info = await client.device_info()
assert info.name == "host-noise-resume"
await asyncio.wait_for(resumed.wait(), timeout=10.0)
assert resumed_count == 1
resumed.clear()
# The resumed session issued a fresh ticket, so it resumes again
await client.disconnect()
await client.connect(login=True)
info = await client.device_info()
assert info.name == "host-noise-resume"
await asyncio.wait_for(resumed.wait(), timeout=10.0)
assert resumed_count == 2
@@ -24,7 +24,6 @@ from .types import (
RunCompiledFunction,
)
pytestmark = pytest.mark.usefixtures("isolated_preferences")
NEW_KEY = PROVISIONING_PSK
@@ -41,15 +41,6 @@ async def _poll_until_exists(path: Path) -> None:
await asyncio.sleep(0.05)
@pytest.fixture(autouse=True)
def isolated_preferences(monkeypatch: pytest.MonkeyPatch, tmp_path) -> Path:
"""Keep host preferences per-test so this test never touches the real
~/.esphome/prefs and never races other tests over ESPHOME_PREFDIR."""
prefdir = tmp_path / "prefs"
monkeypatch.setenv("ESPHOME_PREFDIR", str(prefdir))
return prefdir / f"{DEVICE_NAME}.prefs"
@pytest.mark.asyncio
async def test_host_preferences_suspend_resume(
yaml_config: str,
@@ -58,7 +49,7 @@ async def test_host_preferences_suspend_resume(
isolated_preferences: Path,
) -> None:
"""Test that a running syncer flushes, a suspended one doesn't, and resume restores flushing."""
pref_file = isolated_preferences
pref_file = isolated_preferences / f"{DEVICE_NAME}.prefs"
loop = asyncio.get_running_loop()
saved_in_memory = loop.create_future()
@@ -11,14 +11,6 @@ from .state_utils import InitialStateHelper, require_entity
from .types import APIClientConnectedFactory, RunCompiledFunction
@pytest.fixture(autouse=True)
def isolated_preferences(monkeypatch: pytest.MonkeyPatch, tmp_path) -> None:
"""Keep host preferences per-test so RESTORE_AND_ON never loads a stale value left
behind by a previous run (host preferences otherwise persist to ~/.esphome/prefs,
keyed only by device name)."""
monkeypatch.setenv("ESPHOME_PREFDIR", str(tmp_path / "prefs"))
@pytest.mark.asyncio
async def test_light_initial_state(
yaml_config: str,
+16 -7
View File
@@ -173,6 +173,7 @@ async def test_uart_mock_modbus_no_threshold(
_assert_no_modbus_errors(error_log_lines, warning_log_lines)
@pytest.mark.shared_yaml("uart_mock_modbus_server_injected")
@pytest.mark.asyncio
async def test_uart_mock_modbus_server(
yaml_config: str,
@@ -203,6 +204,7 @@ async def test_uart_mock_modbus_server(
_assert_no_modbus_errors(error_log_lines, warning_log_lines)
@pytest.mark.shared_yaml("uart_mock_modbus_server_injected")
@pytest.mark.asyncio
async def test_uart_mock_modbus_server_read_write(
yaml_config: str,
@@ -231,8 +233,8 @@ async def test_uart_mock_modbus_server_read_write(
"rw_write_1": 4660, # 0x1234 written to reg 0x0001
"rw_read_1": 4660, # reg 0x0001 reads back the just-written value
"rw_read_2": 170, # 0x00AA read from reg 0x0002 in the same request
"rw_write_3": 22136, # 0x5678 written to reg 0x0003
"rw_read_3": 22136, # reg 0x0003 reads back the just-written value
"rw_write_3": 22136, # 0x5678 written to reg 0x0006
"rw_read_3": 22136, # reg 0x0006 reads back the just-written value
}
)
@@ -241,7 +243,8 @@ async def test_uart_mock_modbus_server_read_write(
api_client_connected() as client,
):
await tracker.setup_and_start_scenario(client)
await tracker.await_all(futures)
# The FC 0x17 injections fire last, behind four earlier 100ms delays
await tracker.await_all(futures, timeout=4.0)
_assert_no_modbus_errors(error_log_lines, warning_log_lines)
@@ -296,6 +299,7 @@ async def test_uart_mock_modbus_server_read_write_invalid(
)
@pytest.mark.shared_yaml("uart_mock_modbus_mesh")
@pytest.mark.asyncio
async def test_uart_mock_modbus_server_controller(
yaml_config: str,
@@ -485,6 +489,7 @@ async def test_uart_mock_modbus_server_controller_bits(
_assert_no_modbus_errors(error_log_lines, warning_log_lines)
@pytest.mark.shared_yaml("uart_mock_modbus_mesh")
@pytest.mark.asyncio
async def test_uart_mock_modbus_server_controller_multiple(
yaml_config: str,
@@ -495,7 +500,7 @@ async def test_uart_mock_modbus_server_controller_multiple(
line_callback, error_log_lines, warning_log_lines = _make_modbus_line_callback()
expected_values = {"reg_u_word": 919, "reg_u_word_2": 929}
expected_values = {"multi_reg_a": 919, "multi_reg_b": 929}
tracker = SensorTracker(list(expected_values.keys()))
futures = tracker.expect_all(expected_values)
@@ -706,6 +711,7 @@ async def test_uart_mock_modbus_shared_address(
_assert_no_modbus_errors(error_log_lines, warning_log_lines)
@pytest.mark.shared_yaml("uart_mock_modbus_loopback")
@pytest.mark.asyncio
async def test_uart_mock_modbus_custom_pdu(
yaml_config: str,
@@ -932,6 +938,7 @@ async def test_uart_mock_modbus_broadcast_write(
_assert_no_modbus_errors(error_log_lines, warning_log_lines)
@pytest.mark.shared_yaml("uart_mock_modbus_mesh")
@pytest.mark.asyncio
async def test_uart_mock_modbus_client_read_write(
yaml_config: str,
@@ -947,9 +954,7 @@ async def test_uart_mock_modbus_client_read_write(
"""
line_callback, error_log_lines, warning_log_lines = _make_modbus_line_callback()
tracker = SensorTracker(
["srv_write_1", "srv_read_1", "client_read_0", "client_read_1"]
)
tracker = SensorTracker(["srv_write_1", "client_read_0", "client_read_1"])
futures = tracker.expect_all(
{
"srv_write_1": 4660, # server wrote 0x1234 to reg 0x0001
@@ -967,6 +972,7 @@ async def test_uart_mock_modbus_client_read_write(
_assert_no_modbus_errors(error_log_lines, warning_log_lines)
@pytest.mark.shared_yaml("uart_mock_modbus_loopback")
@pytest.mark.asyncio
async def test_uart_mock_modbus_register_offset(
yaml_config: str,
@@ -1022,6 +1028,7 @@ async def test_uart_mock_modbus_register_offset(
)
@pytest.mark.shared_yaml("uart_mock_modbus_loopback")
@pytest.mark.asyncio
async def test_uart_mock_modbus_lambda_write(
yaml_config: str,
@@ -1058,6 +1065,7 @@ async def test_uart_mock_modbus_lambda_write(
await tracker.await_change(wrote_30, "reg_30", timeout=4.0)
@pytest.mark.shared_yaml("uart_mock_modbus_loopback")
@pytest.mark.asyncio
async def test_uart_mock_modbus_lambda_invert(
yaml_config: str,
@@ -1113,6 +1121,7 @@ async def test_uart_mock_modbus_lambda_invert(
)
@pytest.mark.shared_yaml("uart_mock_modbus_loopback")
@pytest.mark.asyncio
async def test_uart_mock_modbus_deprecated_write_buffer(
yaml_config: str,
+28
View File
@@ -2122,6 +2122,34 @@ def test_get_cpp_changed_components_independent_of_cwd(
) == ["time"]
def test_fixture_map_includes_shared_yaml_markers() -> None:
"""Fixtures named only by shared_yaml markers must map to their test file."""
helpers.get_fixture_to_test_files.cache_clear()
mapping = helpers.get_fixture_to_test_files()
for fixture in (
"uart_mock_modbus_loopback",
"uart_mock_modbus_mesh",
"uart_mock_modbus_server_injected",
):
assert mapping[fixture] == frozenset(
{"tests/integration/test_uart_mock_modbus.py"}
)
def test_no_orphan_integration_fixtures() -> None:
"""Every fixture must reach CI test selection; an orphan selects nothing."""
helpers.get_fixture_to_test_files.cache_clear()
mapping = helpers.get_fixture_to_test_files()
fixtures_dir = (Path(__file__).parent.parent / "integration" / "fixtures").resolve()
fixtures = list(fixtures_dir.glob("*.yaml"))
assert fixtures, f"no fixtures found under {fixtures_dir}"
# cache_init is covered via INTEGRATION_TESTS_TRIGGER_FILES instead
orphans = [
f.stem for f in fixtures if f.stem != "cache_init" and f.stem not in mapping
]
assert not orphans, f"fixtures invisible to CI test selection: {orphans}"
def test_lpt_partition_balances_skewed_weights() -> None:
"""Heavy items spread across groups instead of clustering."""
items = [f"i{n}" for n in range(6)]
@@ -0,0 +1,219 @@
"""Unit tests for script/sync_dependency_versions.py."""
from pathlib import Path
import subprocess
import sys
import pytest
import yamlrocks
sys.path.insert(0, str((Path(__file__).parent / ".." / ".." / "script").resolve()))
import sync_dependency_versions as sync_mod # noqa: E402
PRECOMMIT = """\
# See https://pre-commit.com for more information
repos:
- repo: https://github.com/astral-sh/ruff-pre-commit
# Ruff version.
rev: v0.1.0
hooks:
- id: ruff
- repo: https://github.com/PyCQA/flake8
rev: 7.0.0
hooks:
- id: flake8
- repo: https://github.com/asottile/pyupgrade
rev: v3.0.0
hooks:
- id: pyupgrade
- repo: https://github.com/pre-commit/mirrors-clang-format
rev: v13.0.1
hooks:
- id: clang-format
- repo: https://github.com/adrienverge/yamllint.git
rev: v1.0.0
hooks:
- id: yamllint
- repo: local
hooks:
- id: pylint
"""
REQ_TEST = """\
pylint==4.0.8
flake8==7.1.0
ruff==0.2.0 # comment
pyupgrade==3.0.0
"""
REQ_DEV = """\
clang-format==13.0.1
yamllint==1.0.0
"""
RUFF_REPO = "https://github.com/astral-sh/ruff-pre-commit"
DUPLICATE_RUFF_BLOCK = f" - repo: {RUFF_REPO}\n rev: v0.3.0\n hooks: []\n"
EXPECTED_DRIFT = ["ruff: 0.1.0 -> 0.2.0", "flake8: 7.0.0 -> 7.1.0"]
EXPECTED_PRECOMMIT = PRECOMMIT.replace("rev: v0.1.0", "rev: v0.2.0").replace(
"rev: 7.0.0", "rev: 7.1.0"
)
@pytest.fixture
def root(tmp_path: Path) -> Path:
"""A fake checkout where ruff (v-prefixed) and flake8 (bare) have drifted."""
(tmp_path / ".pre-commit-config.yaml").write_text(PRECOMMIT)
(tmp_path / "requirements_test.txt").write_text(REQ_TEST)
(tmp_path / "requirements_dev.txt").write_text(REQ_DEV)
return tmp_path
def _load(text: str) -> object:
return yamlrocks.loads(text.encode(), option=yamlrocks.OPT_ROUND_TRIP)
@pytest.mark.parametrize(
("requirements", "expected"),
[
("prek==0.5.1 # comment\n", "0.5.1"),
("Prek==0.5.1\n", "0.5.1"),
("other==1.0\nprek==0.5.1\n", "0.5.1"),
("prek>=0.5.1\n", None),
("prek-extra==0.5.1\n", None),
("", None),
],
)
def test_read_requirement_version(requirements: str, expected: str | None) -> None:
assert sync_mod.read_requirement_version(requirements, "prek") == expected
def test_find_repo_entry() -> None:
entry = sync_mod.find_repo_entry(_load(PRECOMMIT), RUFF_REPO)
assert entry["rev"] == "v0.1.0"
@pytest.mark.parametrize(
("text", "message"),
[
("hooks: []\n", "missing key 'repos'"),
("repos:\n - rev: 1.0.0\n", "missing key 'repo'"),
(PRECOMMIT + DUPLICATE_RUFF_BLOCK, "found 2"),
("repos:\n - repo: other\n rev: 1.0.0\n", "found 0"),
],
)
def test_find_repo_entry_errors(text: str, message: str) -> None:
with pytest.raises(sync_mod.SyncError, match=message):
sync_mod.find_repo_entry(_load(text), RUFF_REPO)
@pytest.mark.parametrize(
("rev", "expected"),
[("v0.1.0", ("v", "0.1.0")), ("7.0.0", ("", "7.0.0")), ("'1.0'", ("", "1.0"))],
)
def test_current_rev(rev: str, expected: tuple[str, str]) -> None:
doc = _load(f"repos:\n - repo: {RUFF_REPO}\n rev: {rev}\n")
assert sync_mod.current_rev(doc["repos"][0], RUFF_REPO) == expected
@pytest.mark.parametrize(
("block", "message"),
[(" hooks: []\n", "has no rev"), (" rev: 1.0\n", "not a string: 1.0")],
)
def test_current_rev_errors(block: str, message: str) -> None:
doc = _load(f"repos:\n - repo: {RUFF_REPO}\n{block}")
with pytest.raises(sync_mod.SyncError, match=message):
sync_mod.current_rev(doc["repos"][0], RUFF_REPO)
def test_sync_reports_without_writing(root: Path) -> None:
assert sync_mod.sync(root, write=False) == EXPECTED_DRIFT
assert (root / ".pre-commit-config.yaml").read_text() == PRECOMMIT
def test_sync_writes_keeps_layout_and_is_idempotent(root: Path) -> None:
assert sync_mod.sync(root, write=True) == EXPECTED_DRIFT
assert (root / ".pre-commit-config.yaml").read_text() == EXPECTED_PRECOMMIT
assert sync_mod.sync(root, write=True) == []
def test_sync_does_not_touch_a_config_that_matches(root: Path) -> None:
(root / ".pre-commit-config.yaml").write_text(EXPECTED_PRECOMMIT)
before = (root / ".pre-commit-config.yaml").stat().st_mtime_ns
assert sync_mod.sync(root, write=True) == []
assert (root / ".pre-commit-config.yaml").stat().st_mtime_ns == before
def test_sync_missing_requirement_pin(root: Path) -> None:
(root / "requirements_dev.txt").write_text("")
with pytest.raises(sync_mod.SyncError, match="no 'clang-format==' pin"):
sync_mod.sync(root, write=True)
def test_sync_propagates_config_errors(root: Path) -> None:
(root / ".pre-commit-config.yaml").write_text(PRECOMMIT + DUPLICATE_RUFF_BLOCK)
with pytest.raises(sync_mod.SyncError, match="found 2"):
sync_mod.sync(root, write=True)
def test_main_check_reports_drift(
root: Path, capsys: pytest.CaptureFixture[str]
) -> None:
assert sync_mod.main(["--check", "--root", str(root)]) == 1
assert capsys.readouterr().out.splitlines() == EXPECTED_DRIFT
assert (root / ".pre-commit-config.yaml").read_text() == PRECOMMIT
def test_main_writes_then_check_is_clean(
root: Path, capsys: pytest.CaptureFixture[str]
) -> None:
assert sync_mod.main(["--root", str(root)]) == 0
assert capsys.readouterr().out.splitlines() == EXPECTED_DRIFT
assert sync_mod.main(["--check", "--root", str(root)]) == 0
assert capsys.readouterr().out == ""
def test_main_reports_sync_error(
root: Path, capsys: pytest.CaptureFixture[str]
) -> None:
(root / "requirements_dev.txt").write_text("")
assert sync_mod.main(["--root", str(root)]) == 1
assert (
"error: requirements_dev.txt: no 'clang-format==' pin"
in capsys.readouterr().err
)
def test_main_defaults_to_repo_root(monkeypatch: pytest.MonkeyPatch) -> None:
seen: dict[str, object] = {}
def fake_sync(root: Path, *, write: bool) -> list[str]:
seen["root"] = root
seen["write"] = write
return []
monkeypatch.setattr(sync_mod, "sync", fake_sync)
assert sync_mod.main([]) == 0
assert seen == {"root": sync_mod.REPO_ROOT, "write": True}
def test_repository_is_in_sync() -> None:
"""The real checkout must match; a failure here means a rev has drifted.
Also proves every SYNC_TARGETS entry still resolves in the real files.
"""
assert sync_mod.sync(sync_mod.REPO_ROOT, write=False) == []
def test_cli_entry_point(root: Path) -> None:
"""Run the script the way the workflow does, as a subprocess."""
script = Path(sync_mod.__file__)
result = subprocess.run(
[sys.executable, str(script), "--check", "--root", str(root)],
capture_output=True,
text=True,
check=False,
)
assert result.returncode == 1
assert result.stdout.splitlines() == EXPECTED_DRIFT
@@ -194,17 +194,17 @@ def test_superseded_device_info_fields_still_declared_in_header() -> None:
def test_superseded_device_info_fields_still_encoded_and_sized() -> None:
"""Each superseded field must still be touched by DeviceInfoResponse's
generated encode_msg() and calc_size_msg(), i.e. it is still put on the wire.
generated encode() and calculate_size(), i.e. it is still put on the wire.
"""
encode_body = _extract_function_body(CPP_TEXT, "DeviceInfoResponse::encode_msg")
size_body = _extract_function_body(CPP_TEXT, "DeviceInfoResponse::calc_size_msg")
encode_body = _extract_function_body(CPP_TEXT, "DeviceInfoResponse::encode")
size_body = _extract_function_body(CPP_TEXT, "DeviceInfoResponse::calculate_size")
for field_name in SUPERSEDED_FIELDS:
assert f"msg.{field_name}" in encode_body, (
f"DeviceInfoResponse::encode_msg() no longer references {field_name}. "
assert f"this->{field_name}" in encode_body, (
f"DeviceInfoResponse::encode() no longer references {field_name}. "
f"{DEPRECATED_FIELD_TRAP}"
)
assert f"msg.{field_name}" in size_body, (
f"DeviceInfoResponse::calc_size_msg() no longer references "
assert f"this->{field_name}" in size_body, (
f"DeviceInfoResponse::calculate_size() no longer references "
f"{field_name}. {DEPRECATED_FIELD_TRAP}"
)
@@ -380,13 +380,3 @@ def test_api_version_minor_is_at_least_15() -> None:
"clients to see api_version >= 1.15 in HelloResponse before they will "
"ever request it."
)
def test_generated_encode_calls_keep_the_cursor() -> None:
"""No generated ProtoEncode call may drop the returned cursor."""
dropped = [
line
for line in CPP_TEXT.splitlines()
if "ProtoEncode::" in line and "pos = ProtoEncode::" not in line
]
assert not dropped, dropped[:5]
@@ -15,13 +15,9 @@ import pytest
sys.path.insert(0, str(Path(__file__).parents[4] / "script" / "api_protobuf"))
import aioesphomeapi.api_options_pb2 as pb # noqa: E402
from api_protobuf import ( # noqa: E402
MAX_MESSAGE_ID,
SOURCE_CLIENT,
_make_ifdef_line,
build_message_type,
create_field_type_info,
get_varint64_ifdef,
validate_message_id,
)
@@ -38,26 +34,16 @@ def _file_with_messages(
file_desc = descriptor_pb2.FileDescriptorProto(name="test.proto")
for name, field_type, deprecated in messages:
msg = file_desc.message_type.add(name=name)
field = msg.field.add()
field.CopyFrom(_field(field_type))
field = msg.field.add(name="value", number=1, type=field_type)
field.options.deprecated = deprecated
return file_desc
UINT64 = descriptor_pb2.FieldDescriptorProto.TYPE_UINT64
MESSAGE = descriptor_pb2.FieldDescriptorProto.TYPE_MESSAGE
DOUBLE = descriptor_pb2.FieldDescriptorProto.TYPE_DOUBLE
INT64 = descriptor_pb2.FieldDescriptorProto.TYPE_INT64
SINT64 = descriptor_pb2.FieldDescriptorProto.TYPE_SINT64
UINT32 = descriptor_pb2.FieldDescriptorProto.TYPE_UINT32
INT32 = descriptor_pb2.FieldDescriptorProto.TYPE_INT32
SINT32 = descriptor_pb2.FieldDescriptorProto.TYPE_SINT32
FIXED64 = descriptor_pb2.FieldDescriptorProto.TYPE_FIXED64
FIXED32 = descriptor_pb2.FieldDescriptorProto.TYPE_FIXED32
FLOAT = descriptor_pb2.FieldDescriptorProto.TYPE_FLOAT
BOOL = descriptor_pb2.FieldDescriptorProto.TYPE_BOOL
STRING = descriptor_pb2.FieldDescriptorProto.TYPE_STRING
BYTES = descriptor_pb2.FieldDescriptorProto.TYPE_BYTES
def test_no_varint64_fields() -> None:
@@ -121,190 +107,3 @@ def test_message_id_at_maximum_is_accepted() -> None:
def test_message_id_above_maximum_is_rejected() -> None:
with pytest.raises(ValueError, match="exceeds the plaintext"):
validate_message_id(MAX_MESSAGE_ID + 1, "TooBigMessage")
def _field(
field_type: int, number: int = 1, *, force: bool = False, repeated: bool = False
) -> descriptor_pb2.FieldDescriptorProto:
field = descriptor_pb2.FieldDescriptorProto(
name="value", number=number, type=field_type
)
if repeated:
field.label = descriptor_pb2.FieldDescriptorProto.LABEL_REPEATED
if force:
field.options.Extensions[pb.force] = True
return field
def _encode_field(
field_type: int, number: int = 1, force: bool = False, repeated: bool = False
) -> str:
"""Return the encode statement the generator emits for one encode-only field."""
field = _field(field_type, number, force=force, repeated=repeated)
return create_field_type_info(
field, needs_decode=False, needs_encode=True
).encode_content
SCALAR_TYPES = [
BOOL,
UINT32,
INT32,
UINT64,
INT64,
SINT32,
FLOAT,
FIXED32,
STRING,
BYTES,
]
@pytest.mark.parametrize("field_type", SCALAR_TYPES)
def test_forced_fields_use_the_force_overload_or_raw_writes(field_type: int) -> None:
content = _encode_field(field_type, force=True)
assert (
"_force(" in content
or "write_raw_byte(" in content
or "write_tag_and_fixed32(" in content
), content
@pytest.mark.parametrize("field_type", [FLOAT, FIXED32])
def test_single_byte_tag_fixed32_shares_the_outlined_writer(field_type: int) -> None:
unconditional = _encode_field(field_type, force=True)
assert unconditional.count("write_tag_and_fixed32(pos, 13,") == 1, unconditional
guarded = _encode_field(field_type, force=False)
assert guarded.startswith("if ("), guarded
assert "[[likely]]" in guarded
assert "write_tag_and_fixed32(pos, 13," in guarded
@pytest.mark.parametrize("field_type", [FLOAT, FIXED32])
def test_multi_byte_tag_fixed32_falls_back_to_the_generic_helper(
field_type: int,
) -> None:
content = _encode_field(field_type, number=16)
assert "write_tag_and_fixed32" not in content, content
assert content.startswith("pos = ProtoEncode::encode_"), content
def _decode_case(field_type: int, number: int, *, repeated: bool = False) -> str:
"""Return the decode_field() case the generator emits for one decoded field."""
field = _field(field_type, number, repeated=repeated)
if field_type == MESSAGE:
field.type_name = ".Sub"
return create_field_type_info(
field, needs_decode=True, needs_encode=False
).decode_content
@pytest.mark.parametrize(
("needs_decode", "force", "member"),
[
(False, False, "StringRef value{nullptr, 0}; // null until set, encode only"),
(True, False, "StringRef value{};"),
(False, True, "StringRef value{};"),
],
)
def test_string_fields_default_to_null_only_when_never_read(
needs_decode: bool, force: bool, member: str
) -> None:
"""Only a string that is neither decoded nor force encoded may start as a null StringRef."""
ti = create_field_type_info(
_field(STRING, force=force), needs_decode=needs_decode, needs_encode=True
)
assert ti.public_content == [member]
@pytest.mark.parametrize(
("field_type", "number", "wire_type", "accessor"),
[
(UINT32, 2, "WIRE_TYPE_VARINT", "value.as_varint()"),
(BOOL, 3, "WIRE_TYPE_VARINT", "value.as_bool()"),
(STRING, 1, "WIRE_TYPE_LENGTH_DELIMITED", "value.data()"),
(FLOAT, 4, "WIRE_TYPE_FIXED32", "value.as_float()"),
(FIXED32, 5, "WIRE_TYPE_FIXED32", "value.as_fixed32()"),
],
)
def test_decode_cases_carry_field_number_and_wire_type(
field_type: int, number: int, wire_type: str, accessor: str
) -> None:
"""Each decoded field yields one case keyed on its number and declared wire type."""
case = _decode_case(field_type, number)
lines = case.splitlines()
assert lines[0] == f"case proto_tag({number}, {wire_type}):", case
assert accessor in lines[1], case
assert lines[-1].strip() == "break;", case
@pytest.mark.parametrize(
("field_type", "repeated", "wire_type", "store"),
[
(UINT32, True, "WIRE_TYPE_VARINT", "this->value.push_back(value.as_varint());"),
(
STRING,
True,
"WIRE_TYPE_LENGTH_DELIMITED",
"this->value.push_back(value.as_string());",
),
(
MESSAGE,
False,
"WIRE_TYPE_LENGTH_DELIMITED",
"value.decode_to_message(this->value);",
),
(
MESSAGE,
True,
"WIRE_TYPE_LENGTH_DELIMITED",
"value.decode_to_message(this->value.back());",
),
],
)
def test_repeated_and_message_fields_decode_through_the_same_case_shape(
field_type: int, repeated: bool, wire_type: str, store: str
) -> None:
"""Repeated and sub message fields land in the one switch with their own store."""
case = _decode_case(field_type, 7, repeated=repeated)
lines = case.splitlines()
assert lines[0] == f"case proto_tag(7, {wire_type}):", case
assert store in case, case
if field_type == MESSAGE and repeated:
assert "this->value.emplace_back();" in case, case
assert lines[-1].strip() == "break;", case
def test_a_fixed64_field_fails_at_generation_time() -> None:
"""The decode loop has no 64 bit wire type path, so such a field must never reach it silently."""
desc = descriptor_pb2.DescriptorProto(name="Wide")
desc.field.add(name="ratio", number=1, type=DOUBLE)
with pytest.raises(
ValueError, match="64-bit type 'double' .*ratio.* not supported"
):
build_message_type(desc, {}, {"Wide": SOURCE_CLIENT})
def test_message_gets_a_single_decode_field_override() -> None:
"""All wire types of a decoded message land in one decode_field() switch."""
desc = descriptor_pb2.DescriptorProto(name="Mixed")
desc.field.add(name="name", number=1, type=STRING)
desc.field.add(name="count", number=2, type=UINT32)
desc.field.add(name="level", number=3, type=FLOAT)
header, cpp, _ = build_message_type(desc, {}, {"Mixed": SOURCE_CLIENT})
decl = "void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;"
assert header.count(decl) == 1
assert (
cpp.count(
"void Mixed::decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) {"
)
== 1
)
assert "switch (tag) {" in cpp
assert "const ProtoFieldValue value(data, scalar);" in cpp
for number, wire_type in (
(1, "WIRE_TYPE_LENGTH_DELIMITED"),
(2, "WIRE_TYPE_VARINT"),
(3, "WIRE_TYPE_FIXED32"),
):
assert f"case proto_tag({number}, {wire_type}):" in cpp, cpp
@@ -0,0 +1,37 @@
"""Tests for the udp component configuration schema."""
from __future__ import annotations
import pytest
from esphome.components import udp
from esphome.components.packet_transport import (
CONF_BINARY_SENSORS,
CONF_ENCRYPTION,
CONF_PING_PONG_ENABLE,
CONF_PROVIDERS,
CONF_ROLLING_CODE_ENABLE,
CONF_SENSORS,
)
import esphome.config_validation as cv
@pytest.mark.parametrize(
"option",
[
CONF_PROVIDERS,
CONF_ENCRYPTION,
CONF_PING_PONG_ENABLE,
CONF_ROLLING_CODE_ENABLE,
CONF_SENSORS,
CONF_BINARY_SENSORS,
],
)
def test_relocated_option_rejected(option: str) -> None:
"""Options that moved to packet_transport raise a pointing error."""
with pytest.raises(cv.Invalid) as exc_info:
udp.CONFIG_SCHEMA({option: True})
assert (
f"The '{option}' option should now be configured in the 'packet_transport' component"
in str(exc_info.value)
)
@@ -7,6 +7,7 @@ exercised in their own test modules)."""
import json
import logging
from pathlib import Path
from unittest.mock import Mock
import pytest
@@ -228,6 +229,24 @@ def test_resolve_registry_version_raises_without_pkg_file(monkeypatch):
_resolve_registry_version("owner", "pkg", set())
def test_make_registry_client_skips_private_package_probe(monkeypatch):
"""Our client answers the probe locally without patching PlatformIO's class."""
from platformio.account.client import AccountClient
from platformio.registry.client import RegistryClient
pio_probe = RegistryClient.__dict__["allowed_private_packages"]
monkeypatch.setattr(
AccountClient,
"get_account_info",
Mock(side_effect=AssertionError("account probe must not run")),
)
client = lib._make_registry_client().get_registry_client_instance()
assert client.allowed_private_packages() is False
assert RegistryClient.__dict__["allowed_private_packages"] is pio_probe
def _patch_registry_resolve(monkeypatch: pytest.MonkeyPatch) -> None:
"""Stub the registry lookup so tests never touch the network."""
monkeypatch.setattr(

Some files were not shown because too many files have changed in this diff Show More