mirror of
https://github.com/esphome/esphome.git
synced 2026-10-07 19:44:08 +00:00
[espnow] Keep constant send payloads in shared flash tables (#20271)
This commit is contained in:
+27
-1
@@ -21,7 +21,7 @@ from esphome.const import (
|
||||
CONF_TYPE_ID,
|
||||
CONF_UPDATE_INTERVAL,
|
||||
)
|
||||
from esphome.core import CORE, ID, EsphomeError, Lambda
|
||||
from esphome.core import CORE, ID, EsphomeError, HexInt, Lambda
|
||||
from esphome.cpp_generator import (
|
||||
FlashStringLiteral,
|
||||
LambdaExpression,
|
||||
@@ -35,6 +35,32 @@ from esphome.types import ConfigType, SafeExpType
|
||||
from esphome.util import Registry
|
||||
|
||||
|
||||
def progmem_bytes(name: str, data: bytes | list[int]) -> MockObj:
|
||||
"""Shared PROGMEM table for constant bytes; equal payloads share one, empty is nullptr."""
|
||||
if not data:
|
||||
return cg.nullptr
|
||||
return cg.shared_progmem_array(
|
||||
name, cg.uint8, cg.ArrayInitializer(*(HexInt(x) for x in data))
|
||||
)
|
||||
|
||||
|
||||
async def templatable_bytes(
|
||||
value: Any,
|
||||
args: TemplateArgsType,
|
||||
set_template: MockObj,
|
||||
set_static: MockObj,
|
||||
table_name: str,
|
||||
) -> None:
|
||||
"""Set a TemplatableBytes: a lambda via set_template, constant bytes via set_static."""
|
||||
if cg.is_template(value):
|
||||
fn = await cg.templatable(value, args, cg.std_vector.template(cg.uint8))
|
||||
cg.add(set_template(fn))
|
||||
elif len(value) > 0xFFFF:
|
||||
raise EsphomeError(f"Byte payload is {len(value)} bytes; the maximum is 65535")
|
||||
else:
|
||||
cg.add(set_static(progmem_bytes(table_name, value), len(value)))
|
||||
|
||||
|
||||
def maybe_simple_id(*validators):
|
||||
"""Allow a raw ID to be specified in place of a config block.
|
||||
If the value that's being validated is a dictionary, it's passed as-is to the specified validators. Otherwise, it's
|
||||
|
||||
@@ -24,7 +24,6 @@ from esphome.types import ConfigType
|
||||
CODEOWNERS = ["@jesserockz"]
|
||||
AUTO_LOAD = ["network"]
|
||||
|
||||
byte_vector = cg.std_vector.template(cg.uint8)
|
||||
peer_address_t = cg.std_ns.class_("array").template(cg.uint8, 6)
|
||||
|
||||
espnow_ns = cg.esphome_ns.namespace("espnow")
|
||||
@@ -304,11 +303,12 @@ async def send_action(
|
||||
|
||||
await register_peer(var, config, args)
|
||||
|
||||
data = config.get(CONF_DATA, [])
|
||||
data = config[CONF_DATA]
|
||||
if isinstance(data, str):
|
||||
data = list(data.encode())
|
||||
templ = await cg.templatable(data, args, byte_vector, byte_vector)
|
||||
cg.add(var.set_data(templ))
|
||||
await automation.templatable_bytes(
|
||||
data, args, var.set_data_template, var.set_data_static, "espnow_data"
|
||||
)
|
||||
|
||||
cg.add(var.set_wait_for_sent(config[CONF_WAIT_FOR_SENT]))
|
||||
cg.add(var.set_continue_on_error(config[CONF_CONTINUE_ON_ERROR]))
|
||||
|
||||
@@ -11,7 +11,7 @@ namespace esphome::espnow {
|
||||
|
||||
template<typename... Ts> class SendAction final : public Action<Ts...>, public Parented<ESPNowComponent> {
|
||||
TEMPLATABLE_VALUE(peer_address_t, address);
|
||||
TEMPLATABLE_VALUE(std::vector<uint8_t>, data);
|
||||
TEMPLATABLE_BYTES(data)
|
||||
|
||||
public:
|
||||
void add_on_sent(const std::initializer_list<Action<Ts...> *> &actions) {
|
||||
@@ -58,8 +58,10 @@ template<typename... Ts> class SendAction final : public Action<Ts...>, public P
|
||||
}
|
||||
};
|
||||
peer_address_t address = this->address_.value(x...);
|
||||
std::vector<uint8_t> data = this->data_.value(x...);
|
||||
esp_err_t err = this->parent_->send(address.data(), data, send_callback);
|
||||
esp_err_t err = ESP_OK;
|
||||
this->data_.visit(
|
||||
[&](const uint8_t *data, size_t len) { err = this->parent_->send(address.data(), data, len, send_callback); },
|
||||
x...);
|
||||
if (err != ESP_OK) {
|
||||
send_callback(err);
|
||||
} else if (!this->flags_.wait_for_sent) {
|
||||
|
||||
@@ -69,6 +69,70 @@ template<typename T, typename... X> class TemplatableFn {
|
||||
T (*f_)(X...){nullptr};
|
||||
};
|
||||
|
||||
/// Byte payload that is either a stateless lambda or a static table, which may be in PROGMEM.
|
||||
/// 8 bytes on 32-bit; codegen stores constant payloads as shared flash tables.
|
||||
template<typename... Ts> class TemplatableBytes {
|
||||
public:
|
||||
void set_template(std::vector<uint8_t> (*func)(Ts...)) {
|
||||
this->code_.func = func;
|
||||
this->len_ = -1;
|
||||
}
|
||||
void set_static(const uint8_t *data, uint16_t len) {
|
||||
this->code_.data = data;
|
||||
this->len_ = len;
|
||||
}
|
||||
bool is_static() const { return this->len_ >= 0; }
|
||||
/// Only valid when is_static(); may point to PROGMEM, so read it with progmem_memcpy or progmem_read_byte.
|
||||
const uint8_t *data() const { return this->code_.data; }
|
||||
/// Only valid when is_static().
|
||||
size_t size() const { return static_cast<size_t>(this->len_); }
|
||||
std::vector<uint8_t> value(const Ts &...x) const {
|
||||
if (this->len_ < 0)
|
||||
return this->code_.func(x...);
|
||||
return to_vector(this->code_.data, this->size());
|
||||
}
|
||||
/// Calls fn(const uint8_t *data, size_t len) with the payload readable from RAM: a lambda's vector, a static
|
||||
/// table directly, or on ESP8266 a copy of the PROGMEM table (on the stack up to N bytes).
|
||||
template<size_t N = 64, typename F> void visit(F &&fn, const Ts &...x) const {
|
||||
if (this->len_ < 0) {
|
||||
const std::vector<uint8_t> bytes = this->code_.func(x...);
|
||||
fn(bytes.data(), bytes.size());
|
||||
return;
|
||||
}
|
||||
#ifdef USE_ESP8266
|
||||
SmallBufferWithHeapFallback<N> buf(this->size());
|
||||
if (this->len_ != 0)
|
||||
progmem_memcpy(buf.get(), this->code_.data, this->size());
|
||||
fn(buf.get(), this->size());
|
||||
#else
|
||||
fn(this->code_.data, this->size());
|
||||
#endif
|
||||
}
|
||||
|
||||
protected:
|
||||
static std::vector<uint8_t> to_vector(const uint8_t *data, size_t len) {
|
||||
std::vector<uint8_t> out(len);
|
||||
// An empty payload is (nullptr, 0), and memcpy from nullptr is undefined even for zero bytes.
|
||||
if (len != 0)
|
||||
progmem_memcpy(out.data(), data, len); // byte loads from flash fault on ESP8266
|
||||
return out;
|
||||
}
|
||||
|
||||
union {
|
||||
std::vector<uint8_t> (*func)(Ts...);
|
||||
const uint8_t *data;
|
||||
} code_{};
|
||||
int32_t len_{-1}; // -1: lambda, otherwise the length of the static table
|
||||
};
|
||||
|
||||
#define TEMPLATABLE_BYTES(name) \
|
||||
protected: \
|
||||
TemplatableBytes<Ts...> name##_{}; \
|
||||
\
|
||||
public: \
|
||||
void set_##name##_template(std::vector<uint8_t> (*func)(Ts...)) { this->name##_.set_template(func); } \
|
||||
void set_##name##_static(const uint8_t *data, uint16_t len) { this->name##_.set_static(data, len); }
|
||||
|
||||
// Forward declaration for TemplatableValue (string specialization needs it)
|
||||
template<typename T, typename... X> class TemplatableValue;
|
||||
|
||||
|
||||
@@ -0,0 +1,22 @@
|
||||
---
|
||||
esphome:
|
||||
name: test
|
||||
on_boot:
|
||||
then:
|
||||
- espnow.send:
|
||||
address: 11:22:33:44:55:66
|
||||
data: [0x01, 0x02, 0x03]
|
||||
- espnow.broadcast:
|
||||
data: [0x01, 0x02, 0x03]
|
||||
- espnow.broadcast: "OK"
|
||||
- espnow.broadcast:
|
||||
data: !lambda return {0x09};
|
||||
|
||||
esp32:
|
||||
board: esp32dev
|
||||
|
||||
wifi:
|
||||
ssid: test
|
||||
password: password1
|
||||
|
||||
espnow:
|
||||
@@ -0,0 +1,23 @@
|
||||
"""Tests for ESP-NOW constant send payloads in shared flash tables."""
|
||||
|
||||
from collections.abc import Callable
|
||||
from pathlib import Path
|
||||
import re
|
||||
|
||||
|
||||
def test_constant_payloads_share_tables(
|
||||
generate_main: Callable[[str | Path], str],
|
||||
component_config_path: Callable[[str], Path],
|
||||
) -> None:
|
||||
"""Equal payloads share one table; lambdas stay templates."""
|
||||
main_cpp = generate_main(component_config_path("payload_tables.yaml"))
|
||||
|
||||
tables = dict(
|
||||
re.findall(
|
||||
r"static constexpr uint8_t (\w+)\[\] PROGMEM = (\{[^}]*\});", main_cpp
|
||||
)
|
||||
)
|
||||
assert sorted(tables.values()) == sorted(["{0x01, 0x02, 0x03}", "{0x4F, 0x4B}"])
|
||||
shared = next(k for k, v in tables.items() if v == "{0x01, 0x02, 0x03}")
|
||||
assert main_cpp.count(f"set_data_static({shared}, 3);") == 2
|
||||
assert "set_data_template(" in main_cpp
|
||||
@@ -0,0 +1,43 @@
|
||||
#include <gtest/gtest.h>
|
||||
#include <vector>
|
||||
#include "esphome/core/automation.h"
|
||||
|
||||
namespace esphome::testing {
|
||||
|
||||
static const uint8_t PAYLOAD[] = {1, 2, 3};
|
||||
|
||||
static std::vector<uint8_t> repeat(int count) { return std::vector<uint8_t>(count, 7); }
|
||||
|
||||
TEST(TemplatableBytesTest, VisitStaticTable) {
|
||||
TemplatableBytes<> bytes;
|
||||
bytes.set_static(PAYLOAD, sizeof(PAYLOAD));
|
||||
std::vector<uint8_t> seen;
|
||||
bytes.visit([&](const uint8_t *data, size_t len) { seen.assign(data, data + len); });
|
||||
EXPECT_EQ(seen, (std::vector<uint8_t>{1, 2, 3}));
|
||||
}
|
||||
|
||||
TEST(TemplatableBytesTest, VisitWithLargerStackBuffer) {
|
||||
TemplatableBytes<> bytes;
|
||||
bytes.set_static(PAYLOAD, sizeof(PAYLOAD));
|
||||
size_t seen = 0;
|
||||
bytes.visit<256>([&](const uint8_t *, size_t len) { seen = len; });
|
||||
EXPECT_EQ(seen, sizeof(PAYLOAD));
|
||||
}
|
||||
|
||||
TEST(TemplatableBytesTest, VisitEmptyStaticTable) {
|
||||
TemplatableBytes<> bytes;
|
||||
bytes.set_static(nullptr, 0);
|
||||
size_t seen = 1;
|
||||
bytes.visit([&](const uint8_t *, size_t len) { seen = len; });
|
||||
EXPECT_EQ(seen, 0u);
|
||||
}
|
||||
|
||||
TEST(TemplatableBytesTest, VisitLambdaWithArgument) {
|
||||
TemplatableBytes<int> bytes;
|
||||
bytes.set_template(repeat);
|
||||
std::vector<uint8_t> seen;
|
||||
bytes.visit([&](const uint8_t *data, size_t len) { seen.assign(data, data + len); }, 4);
|
||||
EXPECT_EQ(seen, (std::vector<uint8_t>(4, 7)));
|
||||
}
|
||||
|
||||
} // namespace esphome::testing
|
||||
@@ -20,6 +20,7 @@ from esphome.automation import (
|
||||
has_non_synchronous_actions,
|
||||
literal_with_length,
|
||||
maybe_simple_id,
|
||||
progmem_bytes,
|
||||
register_apply_action,
|
||||
register_apply_condition,
|
||||
register_bare_action,
|
||||
@@ -28,6 +29,7 @@ from esphome.automation import (
|
||||
register_parented_condition,
|
||||
register_simple_action,
|
||||
register_simple_condition,
|
||||
templatable_bytes,
|
||||
)
|
||||
import esphome.codegen as cg
|
||||
import esphome.config_validation as cv
|
||||
@@ -989,3 +991,48 @@ async def test_apply_condition_string_lambda_paths(
|
||||
text = _apply_definition(mock_cg)
|
||||
assert expected in text
|
||||
assert ("-> std::string {" in text) is called
|
||||
|
||||
|
||||
def test_progmem_bytes_shares_equal_payloads() -> None:
|
||||
CORE.config = {}
|
||||
a = progmem_bytes("payload", [1, 2])
|
||||
b = progmem_bytes("payload", b"\x01\x02")
|
||||
assert a is b
|
||||
assert str(progmem_bytes("payload", [])) == "nullptr"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_templatable_bytes_static_payload() -> None:
|
||||
CORE.config = {}
|
||||
var = MockObj("act", "->")
|
||||
await templatable_bytes(
|
||||
[0xA1, 0x02], [], var.set_code_template, var.set_code_static, "payload"
|
||||
)
|
||||
assert "act->set_code_static(payload, 2);" in CORE.cpp_main_section
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_templatable_bytes_lambda_payload() -> None:
|
||||
CORE.config = {}
|
||||
var = MockObj("act", "->")
|
||||
await templatable_bytes(
|
||||
Lambda("return {0x01, 0x02};"),
|
||||
[],
|
||||
var.set_code_template,
|
||||
var.set_code_static,
|
||||
"payload",
|
||||
)
|
||||
text = CORE.cpp_main_section
|
||||
assert "act->set_code_template(" in text
|
||||
assert "-> std::vector<uint8_t>" in text
|
||||
assert "set_code_static" not in text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_templatable_bytes_rejects_oversized_payload() -> None:
|
||||
CORE.config = {}
|
||||
var = MockObj("act", "->")
|
||||
with pytest.raises(EsphomeError, match="maximum is 65535"):
|
||||
await templatable_bytes(
|
||||
[0] * 0x10000, [], var.set_code_template, var.set_code_static, "payload"
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user