Compare commits

..
Author SHA1 Message Date
J. Nick Koston 455e1d6374 [api] Assert the switch frame count instead of reading for it separately 2026-09-07 15:57:12 +02:00
J. Nick Koston 3af1d50bce [api] Trim the field free message test and a duplicated generator note 2026-09-07 15:57:12 +02:00
J. Nick Koston 842f354a05 [api] Add an integration test for field free messages
Ping, device info, list entities done and disconnect all travel through
the ProtoMessage static entry points now that the no-op thunk is gone.
2026-09-07 15:57:12 +02:00
J. Nick Koston af9b59d4bd [api] Trim the type erased entry point comments 2026-09-07 15:57:12 +02:00
J. Nick Koston 55fc5a10de [api] Tighten the ProtoMessage default entry point comment 2026-09-07 15:57:12 +02:00
J. Nick Koston 79927b918b [api] Clarify which encode entry points forward on ProtoMessage
The base class defaults are independent no-ops; only generated message
classes forward encode() and calculate_size() to their statics.
2026-09-07 15:57:12 +02:00
J. Nick Koston 1b070629bc [api] Make generated encode and size entry points type erased
Every message sent through send_message or the entity paths needed a
proto_encode_msg<T> thunk (17 bytes on xtensa) and, for entity state
and info messages, a calc_size<T> thunk, because the generated encode
and calculate_size were member functions and the connection code wants
plain function pointers over const void *.

The generator now emits the bodies as static encode_msg(const void *)
and calc_size_msg(const void *) functions, so &T::encode_msg is
already a MessageEncodeFn and the thunks disappear. The member
encode() and calculate_size() remain as inline forwarders for direct
callers. ProtoMessage carries the same static defaults for messages
without fields, which also removes the separate no-op encode thunk.
2026-09-07 15:57:12 +02:00
J. Nick Koston c4e1360cdf [api] Check the encoded end against the reserved size under ESPHOME_DEBUG_API
The fixed32 store helper moves to a private section since it neither bounds checks nor
advances the cursor, its comment describes the path each target takes, the generated
file scan flags any ProtoEncode call that does not assign the cursor, and StateWaiter
timeouts can carry a label so gathered waits are told apart.
2026-09-07 15:57:10 +02:00
J. Nick Koston 7c774699d7 [api] Outline the fixed32 writers on ARM cores without unaligned access too
Cortex-M0+ and ARM9 turn the four byte unaligned store into a memcpy call with a stack
temporary at every fixed32 field, and the outlined helper itself became a memcpy call
there, so the helper now spells out the byte stores. Xtensa and host objects are byte for
byte unchanged; on the RP2040 bench config the api object loses 28 bytes and the fixed32
memcpy calls.
2026-09-07 15:24:30 +02:00
J. Nick Koston ea71a24a9b [api] Mark the last two raw varint writers nodiscard and make StateWaiter failures visible
A predicate that raises now fails its wait instead of dying inside the state callback,
and a timeout names the predicate it was waiting for.
2026-09-07 14:01:34 +02:00
J. Nick Koston 709a1e1eb6 [api] Mark the raw encode helpers nodiscard too and drop a duplicate cursor test
The generated file scan already covers every emitted call, so the parametrized copy of
the same assertion goes.
2026-09-07 13:48:53 +02:00
J. Nick Koston 822b701792 [api] Mark the cursor returning encode helpers nodiscard
A call that drops the returned cursor would silently truncate the message, so the
compiler now warns on it and a unit test scans the generated file for the same mistake.
Also corrects the outlining comment for ESP8266, where the inline write is a few byte
stores rather than one, and the RAW_ENCODE_MAP annotation.
2026-09-07 12:37:48 +02:00
J. Nick Koston d2e4d2c46a [api] Outline the fixed32 writers only where memcpy is a call
On the ESP8266 the inline write was already a single store, so the
outlined helper cost a call per fixed32 field: sensor state encode went
from 615 to 864 ns on a d1 mini. ESP32 builds pass -fno-builtin-memcpy,
where the shared copy is both smaller and faster (562 to 328 ns on an
atom), so the gate is now USE_ESP32.
2026-09-07 11:40:01 +02:00
J. Nick Koston adbbda4072 [api] Emit every encode call through one generator helper
_encode_call() owns the cursor assignment and the _force suffix, so
the convention lives in one place instead of at every emission site;
the fixed32 fast path is an arm of the generic encode_content keyed by
a per type value template. write_fixed32_le uses convert_little_endian
instead of its own byte order switch. The integration test shares a
StateWaiter from state_utils and leaves the disconnect to the fixture.
2026-09-07 11:09:07 +02:00
J. Nick Koston 252bf6ea6a [api] Add an integration test for the encode branch boundaries
Covers a zero float that is skipped on the wire, a fixed32 state, a
negative int32, list entity strings and text states whose length
prefix needs two varint bytes, a two byte field tag through the
device info area, and the field free disconnect exchange.
2026-09-07 10:58:31 +02:00
J. Nick Koston 8ec9305688 [api] Trim the encode helper comments 2026-09-07 10:48:18 +02:00
J. Nick Koston 490aca17e6 [api] Share the fixed32 emission between float and fixed32 fields
One helper next to the other precomputed tag paths decides how a
single byte tag fixed32 field is written; the float and fixed32 types
only differ in the value expression. Drop the non forced std::string
encode_string overload, which the generator never emits, and build the
generator tests from one block of field type constants.
2026-09-07 10:33:21 +02:00
J. Nick Koston b77e2441d4 [api] Undefine PROTO_OUTLINE_FOR_SIZE after the encode helpers
The macro only exists for the two fixed32 writers in ProtoEncode, so
drop it once the class is complete instead of leaking it into every
translation unit that includes proto.h.
2026-09-07 10:09:46 +02:00
J. Nick Koston 3b14f4dfc8 [api] Pass the encode cursor by value through the protobuf helpers
The ProtoEncode helpers took the write cursor by reference and a
bool force flag. At -Os the compiler outlines most of them, so every
call site had to keep pos in a stack slot and pass its address, plus
a constant for the flag. The helpers now take the cursor by value and
return the advanced cursor, so consecutive calls chain through the
return register; forced fields call a _force overload instead of
passing a flag.

The fixed32 writers use __builtin_memcpy, which stays a builtin under
ESP-IDF's -fno-builtin-memcpy, and are outlined on embedded targets so
each fixed32 or float field is a short call instead of an inline
memcpy call. Non-forced float and fixed32 fields with a single-byte
tag share the same writer behind a zero check.

Generated encode bodies shrink by 18 percent on an ESP32 IDF proxy
build (2360 to 1932 bytes for 27 messages); entity messages gain the
most, for example ListEntitiesSensorResponse::encode 190 to 134 bytes
and SensorStateResponse::encode 78 to 49 bytes.
2026-09-07 09:21:54 +02:00
18 changed files with 2913 additions and 1942 deletions
+6 -1
View File
@@ -2255,7 +2255,12 @@ bool APIConnection::send_message_(uint32_t payload_size, uint16_t message_type,
// Capacity reserved above, cannot fail // Capacity reserved above, cannot fail
(void) shared_buf.resize(write_start + payload_size); (void) shared_buf.resize(write_start + payload_size);
ProtoWriteBuffer buffer{&shared_buf, write_start}; ProtoWriteBuffer buffer{&shared_buf, write_start};
encode_fn(msg, buffer PROTO_ENCODE_DEBUG_INIT(&shared_buf)); 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
return this->send_buffer(ProtoWriteBuffer{&shared_buf}, message_type); return this->send_buffer(ProtoWriteBuffer{&shared_buf}, message_type);
} }
// encode_to_buffer is defined inline in api_connection.h (ESPHOME_ALWAYS_INLINE) // encode_to_buffer is defined inline in api_connection.h (ESPHOME_ALWAYS_INLINE)
+6 -24
View File
@@ -345,11 +345,7 @@ class APIConnection final : public APIServerConnectionBase {
/// Returns false as soon as the TCP buffer is full. Marked nodiscard so we /// 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. /// have no silent failures: every caller must handle (or log) a refusal.
template<typename T> [[nodiscard]] bool send_message(const T &msg) { template<typename T> [[nodiscard]] bool send_message(const T &msg) {
if constexpr (T::ESTIMATED_SIZE == 0) { return this->send_message_(T::calc_size_msg(&msg), T::MESSAGE_TYPE, &T::encode_msg, &msg);
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. /// Clear the shared write buffer and reserve space for the first message.
@@ -405,16 +401,6 @@ class APIConnection final : public APIServerConnectionBase {
void process_state_subscriptions_(); void process_state_subscriptions_();
#endif #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 // Non-template buffer management for send_message
bool send_message_(uint32_t payload_size, uint16_t message_type, MessageEncodeFn encode_fn, const void *msg); bool send_message_(uint32_t payload_size, uint16_t message_type, MessageEncodeFn encode_fn, const void *msg);
@@ -433,11 +419,7 @@ class APIConnection final : public APIServerConnectionBase {
// Hot paths (state/info) go through fill_and_encode_entity_state/info instead. // 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. // 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) { template<typename T> static uint16_t encode_message_to_buffer(T &msg, APIConnection *conn, uint32_t remaining_size) {
if constexpr (T::ESTIMATED_SIZE == 0) { return encode_to_buffer_slow(T::calc_size_msg(&msg), &T::encode_msg, &msg, conn, remaining_size);
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 // Non-template core — fills state fields and encodes
@@ -449,7 +431,7 @@ class APIConnection final : public APIServerConnectionBase {
template<typename T> template<typename T>
static uint16_t fill_and_encode_entity_state(EntityBase *entity, T &msg, APIConnection *conn, static uint16_t fill_and_encode_entity_state(EntityBase *entity, T &msg, APIConnection *conn,
uint32_t remaining_size) { uint32_t remaining_size) {
return fill_and_encode_entity_state(entity, msg, &calc_size<T>, &proto_encode_msg<T>, conn, remaining_size); return fill_and_encode_entity_state(entity, msg, &T::calc_size_msg, &T::encode_msg, conn, remaining_size);
} }
// Non-template core — fills info fields, allocates buffers, and encodes // Non-template core — fills info fields, allocates buffers, and encodes
@@ -461,7 +443,7 @@ class APIConnection final : public APIServerConnectionBase {
template<typename T> template<typename T>
static uint16_t fill_and_encode_entity_info(EntityBase *entity, T &msg, APIConnection *conn, static uint16_t fill_and_encode_entity_info(EntityBase *entity, T &msg, APIConnection *conn,
uint32_t remaining_size) { uint32_t remaining_size) {
return fill_and_encode_entity_info(entity, msg, &calc_size<T>, &proto_encode_msg<T>, conn, remaining_size); return fill_and_encode_entity_info(entity, msg, &T::calc_size_msg, &T::encode_msg, conn, remaining_size);
} }
// Non-template core — fills device_class, then delegates to fill_and_encode_entity_info // Non-template core — fills device_class, then delegates to fill_and_encode_entity_info
@@ -475,8 +457,8 @@ class APIConnection final : public APIServerConnectionBase {
static uint16_t fill_and_encode_entity_info_with_device_class(EntityBase *entity, T &msg, static uint16_t fill_and_encode_entity_info_with_device_class(EntityBase *entity, T &msg,
StringRef &device_class_field, APIConnection *conn, StringRef &device_class_field, APIConnection *conn,
uint32_t remaining_size) { uint32_t remaining_size) {
return fill_and_encode_entity_info_with_device_class(entity, msg, device_class_field, &calc_size<T>, return fill_and_encode_entity_info_with_device_class(entity, msg, device_class_field, &T::calc_size_msg,
&proto_encode_msg<T>, conn, remaining_size); &T::encode_msg, conn, remaining_size);
} }
#ifdef USE_VOICE_ASSISTANT #ifdef USE_VOICE_ASSISTANT
@@ -46,7 +46,13 @@ inline uint16_t ESPHOME_ALWAYS_INLINE APIConnection::encode_to_buffer(uint32_t c
return 0; return 0;
} }
ProtoWriteBuffer buffer{&shared_buf, shared_buf.size() - calculated_size}; ProtoWriteBuffer buffer{&shared_buf, shared_buf.size() - calculated_size};
encode_fn(msg, buffer PROTO_ENCODE_DEBUG_INIT(&shared_buf)); 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
return total_calculated_size; return total_calculated_size;
} }
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+198 -134
View File
@@ -287,19 +287,31 @@ class ProtoWriteBuffer {
uint8_t *pos_; 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. // 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_1_BYTE = 1 << 7; // 128
constexpr uint32_t VARINT_MAX_2_BYTE = 1 << 14; // 16384 constexpr uint32_t VARINT_MAX_2_BYTE = 1 << 14; // 16384
/// Static encode helpers for generated encode() functions. /// Static encode helpers for the generated encode bodies. Each takes the write cursor by value and
/// Generated code hoists buffer.pos_ into a local uint8_t *__restrict__ pos, /// returns it advanced, so outlined calls at -Os chain through the return register instead of a
/// then calls these methods which take pos by reference. No struct, no overhead. /// stack slot. Helpers without a _force suffix skip fields holding the proto3 default.
/// For sub-messages, pos is synced back to buffer before the call and reloaded after.
class ProtoEncode { class ProtoEncode {
public: public:
/// Write a multi-byte varint directly through a pos pointer. /// Write a multi-byte varint directly through a pos pointer.
template<typename T> template<typename T>
static inline void encode_varint_raw_loop(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, T value) { [[nodiscard]] static inline uint8_t *encode_varint_raw_loop(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
T value) {
do { do {
PROTO_ENCODE_CHECK_BOUNDS(pos, 1); PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
*pos++ = static_cast<uint8_t>(value | 0x80); *pos++ = static_cast<uint8_t>(value | 0x80);
@@ -307,48 +319,49 @@ class ProtoEncode {
} while (value > 0x7F); } while (value > 0x7F);
PROTO_ENCODE_CHECK_BOUNDS(pos, 1); PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
*pos++ = static_cast<uint8_t>(value); *pos++ = static_cast<uint8_t>(value);
return pos;
} }
static inline void ESPHOME_ALWAYS_INLINE encode_varint_raw(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, [[nodiscard]] static inline uint8_t *ESPHOME_ALWAYS_INLINE
uint32_t value) { encode_varint_raw(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint32_t value) {
if (value < VARINT_MAX_1_BYTE) [[likely]] { if (value < VARINT_MAX_1_BYTE) [[likely]] {
PROTO_ENCODE_CHECK_BOUNDS(pos, 1); PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
*pos++ = static_cast<uint8_t>(value); *pos++ = static_cast<uint8_t>(value);
return; return pos;
} }
encode_varint_raw_loop(pos PROTO_ENCODE_DEBUG_ARG, value); return 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). /// Encode a varint that is expected to be 1-2 bytes (e.g. zigzag RSSI, small lengths).
static inline void ESPHOME_ALWAYS_INLINE encode_varint_raw_short(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, [[nodiscard]] static inline uint8_t *ESPHOME_ALWAYS_INLINE
uint32_t value) { encode_varint_raw_short(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint32_t value) {
if (value < VARINT_MAX_1_BYTE) [[likely]] { if (value < VARINT_MAX_1_BYTE) [[likely]] {
PROTO_ENCODE_CHECK_BOUNDS(pos, 1); PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
*pos++ = static_cast<uint8_t>(value); *pos++ = static_cast<uint8_t>(value);
return; return pos;
} }
if (value < VARINT_MAX_2_BYTE) [[likely]] { if (value < VARINT_MAX_2_BYTE) [[likely]] {
PROTO_ENCODE_CHECK_BOUNDS(pos, 2); PROTO_ENCODE_CHECK_BOUNDS(pos, 2);
*pos++ = static_cast<uint8_t>(value | 0x80); *pos++ = static_cast<uint8_t>(value | 0x80);
*pos++ = static_cast<uint8_t>(value >> 7); *pos++ = static_cast<uint8_t>(value >> 7);
return; return pos;
} }
encode_varint_raw_loop(pos PROTO_ENCODE_DEBUG_ARG, value); return encode_varint_raw_loop(pos PROTO_ENCODE_DEBUG_ARG, value);
} }
static inline void ESPHOME_ALWAYS_INLINE encode_varint_raw_64(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, [[nodiscard]] static inline uint8_t *ESPHOME_ALWAYS_INLINE
uint64_t value) { encode_varint_raw_64(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint64_t value) {
if (value < VARINT_MAX_1_BYTE) [[likely]] { if (value < VARINT_MAX_1_BYTE) [[likely]] {
PROTO_ENCODE_CHECK_BOUNDS(pos, 1); PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
*pos++ = static_cast<uint8_t>(value); *pos++ = static_cast<uint8_t>(value);
return; return pos;
} }
encode_varint_raw_loop(pos PROTO_ENCODE_DEBUG_ARG, value); return encode_varint_raw_loop(pos PROTO_ENCODE_DEBUG_ARG, value);
} }
/// Encode a 48-bit MAC address (stored in a uint64) as varint. /// 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 /// 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 /// 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. /// with no per-byte branch. Falls back to the general loop otherwise.
/// Caller must guarantee value fits in 48 bits (checked in debug builds). /// Caller must guarantee value fits in 48 bits (checked in debug builds).
static inline void ESPHOME_ALWAYS_INLINE encode_varint_raw_48bit(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, [[nodiscard]] static inline uint8_t *ESPHOME_ALWAYS_INLINE
uint64_t value) { encode_varint_raw_48bit(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint64_t value) {
#ifdef ESPHOME_DEBUG_API #ifdef ESPHOME_DEBUG_API
assert(value < (1ULL << (MAC_ADDRESS_SIZE * 8)) && "encode_varint_raw_48bit: value exceeds 48 bits"); assert(value < (1ULL << (MAC_ADDRESS_SIZE * 8)) && "encode_varint_raw_48bit: value exceeds 48 bits");
#endif #endif
@@ -363,38 +376,39 @@ class ProtoEncode {
pos[4] = static_cast<uint8_t>((value >> 28) | 0x80); pos[4] = static_cast<uint8_t>((value >> 28) | 0x80);
pos[5] = static_cast<uint8_t>((value >> 35) | 0x80); pos[5] = static_cast<uint8_t>((value >> 35) | 0x80);
pos[6] = static_cast<uint8_t>(value >> 42); pos[6] = static_cast<uint8_t>(value >> 42);
pos += 7; return pos + 7;
return;
} }
encode_varint_raw_64(pos PROTO_ENCODE_DEBUG_ARG, value); return encode_varint_raw_64(pos PROTO_ENCODE_DEBUG_ARG, value);
} }
static inline void ESPHOME_ALWAYS_INLINE encode_field_raw(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, [[nodiscard]] static inline uint8_t *ESPHOME_ALWAYS_INLINE
uint32_t field_id, uint32_t type) { 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); return encode_varint_raw(pos PROTO_ENCODE_DEBUG_ARG, (field_id << 3) | type);
} }
/// Write a single precomputed tag byte. Tag must be < 128. /// Write a single precomputed tag byte. Tag must be < 128.
static inline void ESPHOME_ALWAYS_INLINE write_raw_byte(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, [[nodiscard]] static inline uint8_t *ESPHOME_ALWAYS_INLINE
uint8_t b) { write_raw_byte(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint8_t b) {
PROTO_ENCODE_CHECK_BOUNDS(pos, 1); PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
*pos++ = b; *pos++ = b;
return pos;
} }
/// Reserve one byte for later backpatch (e.g., sub-message length). /// Reserve one byte for later backpatch (e.g., sub-message length).
/// Advances pos past the reserved byte without writing a value. /// Advances pos past the reserved byte without writing a value.
static inline void ESPHOME_ALWAYS_INLINE reserve_byte(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM) { [[nodiscard]] static inline uint8_t *ESPHOME_ALWAYS_INLINE
reserve_byte(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM) {
PROTO_ENCODE_CHECK_BOUNDS(pos, 1); PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
pos++; return pos + 1;
} }
/// Write raw bytes to the buffer (no tag, no length prefix). /// Write raw bytes to the buffer (no tag, no length prefix).
static inline void ESPHOME_ALWAYS_INLINE encode_raw(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, [[nodiscard]] static inline uint8_t *ESPHOME_ALWAYS_INLINE
const void *data, size_t len) { encode_raw(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, const void *data, size_t len) {
PROTO_ENCODE_CHECK_BOUNDS(pos, len); PROTO_ENCODE_CHECK_BOUNDS(pos, len);
std::memcpy(pos, data, len); std::memcpy(pos, data, len);
pos += len; return pos + len;
} }
/// Encode tag + 1-byte length + raw string data. For strings with max_data_length < 128. /// 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). /// Tag must be a single-byte varint (< 128). Always encodes (no zero check).
static inline void encode_short_string_force(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint8_t tag, [[nodiscard]] static inline uint8_t *encode_short_string_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
const StringRef &ref) { uint8_t tag, const StringRef &ref) {
#ifdef ESPHOME_DEBUG_API #ifdef ESPHOME_DEBUG_API
assert(ref.size() < 128 && "encode_short_string_force: string exceeds max_data_length < 128"); assert(ref.size() < 128 && "encode_short_string_force: string exceeds max_data_length < 128");
#endif #endif
@@ -402,137 +416,191 @@ class ProtoEncode {
pos[0] = tag; pos[0] = tag;
pos[1] = static_cast<uint8_t>(ref.size()); pos[1] = static_cast<uint8_t>(ref.size());
std::memcpy(pos + 2, ref.c_str(), ref.size()); std::memcpy(pos + 2, ref.c_str(), ref.size());
pos += 2 + ref.size(); return pos + 2 + ref.size();
} }
/// Write a precomputed tag byte + 32-bit value in one operation. /// Write a precomputed tag byte + 32-bit value. Outlined on embedded: one copy beats inline stores per field.
static inline void ESPHOME_ALWAYS_INLINE write_tag_and_fixed32(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, [[nodiscard]] static PROTO_OUTLINE_FOR_SIZE uint8_t *write_tag_and_fixed32(
uint8_t tag, uint32_t value) { uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint8_t tag, uint32_t value) {
PROTO_ENCODE_CHECK_BOUNDS(pos, 5); PROTO_ENCODE_CHECK_BOUNDS(pos, 5);
pos[0] = tag; pos[0] = tag;
#if __BYTE_ORDER__ == __ORDER_LITTLE_ENDIAN__ write_fixed32_le(pos + 1, value);
std::memcpy(pos + 1, &value, 4); return pos + 5;
#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;
} }
static inline void encode_string(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id, [[nodiscard]] static inline uint8_t *encode_string_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
const char *string, size_t len, bool force = false) { uint32_t field_id, const char *string, size_t len) {
if (len == 0 && !force) pos = encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 2); // type 2: Length-delimited string
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 // NOLINTNEXTLINE(readability-inconsistent-ifelse-braces) -- false positive on [[likely]] attribute
if (len < VARINT_MAX_1_BYTE) [[likely]] { if (len < VARINT_MAX_1_BYTE) [[likely]] {
PROTO_ENCODE_CHECK_BOUNDS(pos, 1 + len); PROTO_ENCODE_CHECK_BOUNDS(pos, 1 + len);
*pos++ = static_cast<uint8_t>(len); *pos++ = static_cast<uint8_t>(len);
} else { } else {
encode_varint_raw_loop(pos PROTO_ENCODE_DEBUG_ARG, len); pos = encode_varint_raw_loop(pos PROTO_ENCODE_DEBUG_ARG, len);
PROTO_ENCODE_CHECK_BOUNDS(pos, len); PROTO_ENCODE_CHECK_BOUNDS(pos, len);
} }
std::memcpy(pos, string, len); std::memcpy(pos, string, len);
pos += len; return pos + len;
} }
static inline void encode_string(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id, [[nodiscard]] static inline uint8_t *encode_string(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
const std::string &value, bool force = false) { uint32_t field_id, const char *string, size_t len) {
encode_string(pos PROTO_ENCODE_DEBUG_ARG, field_id, value.data(), value.size(), force); 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, [[nodiscard]] static inline uint8_t *encode_string_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
const StringRef &ref, bool force = false) { uint32_t field_id, const std::string &value) {
encode_string(pos PROTO_ENCODE_DEBUG_ARG, field_id, ref.c_str(), ref.size(), force); return encode_string_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, value.data(), value.size());
} }
static inline void encode_bytes(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id, [[nodiscard]] static inline uint8_t *encode_string(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
const uint8_t *data, size_t len, bool force = false) { uint32_t field_id, const StringRef &ref) {
encode_string(pos PROTO_ENCODE_DEBUG_ARG, field_id, reinterpret_cast<const char *>(data), len, force); return encode_string(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, [[nodiscard]] static inline uint8_t *encode_string_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t value, bool force = false) { uint32_t field_id, const StringRef &ref) {
if (value == 0 && !force) return encode_string_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, ref.c_str(), ref.size());
return;
encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 0);
encode_varint_raw(pos PROTO_ENCODE_DEBUG_ARG, value);
} }
static inline void encode_uint64(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id, [[nodiscard]] static inline uint8_t *encode_bytes(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint64_t value, bool force = false) { uint32_t field_id, const uint8_t *data, size_t len) {
if (value == 0 && !force) return encode_string(pos PROTO_ENCODE_DEBUG_ARG, field_id, reinterpret_cast<const char *>(data), len);
return;
encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 0);
encode_varint_raw_64(pos PROTO_ENCODE_DEBUG_ARG, value);
} }
static inline void encode_bool(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id, bool value, [[nodiscard]] static inline uint8_t *encode_bytes_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
bool force = false) { uint32_t field_id, const uint8_t *data, size_t len) {
if (!value && !force) return encode_string_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, reinterpret_cast<const char *>(data), len);
return; }
encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 0); [[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);
PROTO_ENCODE_CHECK_BOUNDS(pos, 1); PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
*pos++ = value ? 0x01 : 0x00; *pos++ = value ? 0x01 : 0x00;
return pos;
} }
static inline void encode_fixed32(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id, [[nodiscard]] static inline uint8_t *encode_bool(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t value, bool force = false) { uint32_t field_id, bool value) {
if (value == 0 && !force) if (!value)
return; return pos;
encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 5); 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);
PROTO_ENCODE_CHECK_BOUNDS(pos, 4); PROTO_ENCODE_CHECK_BOUNDS(pos, 4);
#if __BYTE_ORDER__ == __ORDER_LITTLE_ENDIAN__ write_fixed32_le(pos, value);
std::memcpy(pos, &value, 4); return pos + 4;
pos += 4; }
#else [[nodiscard]] static inline uint8_t *encode_fixed32(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
*pos++ = (value >> 0) & 0xFF; uint32_t field_id, uint32_t value) {
*pos++ = (value >> 8) & 0xFF; if (value == 0)
*pos++ = (value >> 16) & 0xFF; return pos;
*pos++ = (value >> 24) & 0xFF; return encode_fixed32_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, value);
#endif
} }
// NOTE: Wire type 1 (64-bit fixed: double, fixed64, sfixed64) is intentionally // NOTE: Wire type 1 (64-bit fixed: double, fixed64, sfixed64) is intentionally
// not supported to reduce overhead on embedded systems. All ESPHome devices are // 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 // 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. // is needed in the future, the necessary encoding/decoding functions must be added.
static inline void encode_float(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id, float value, [[nodiscard]] static inline uint8_t *encode_float(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
bool force = false) { uint32_t field_id, float value) {
uint32_t raw = float_to_raw(value); return encode_fixed32(pos PROTO_ENCODE_DEBUG_ARG, field_id, float_to_raw(value));
if (raw == 0 && !force)
return;
encode_fixed32(pos PROTO_ENCODE_DEBUG_ARG, field_id, raw);
} }
static inline void encode_int32(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id, int32_t value, [[nodiscard]] static inline uint8_t *encode_float_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
bool force = false) { 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) {
if (value < 0) { if (value < 0) {
// negative int32 is always 10 byte long // negative int32 is always 10 byte long
encode_uint64(pos PROTO_ENCODE_DEBUG_ARG, field_id, static_cast<uint64_t>(value), force); return encode_uint64_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, static_cast<uint64_t>(value));
return;
} }
encode_uint32(pos PROTO_ENCODE_DEBUG_ARG, field_id, static_cast<uint32_t>(value), force); return encode_uint32_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, static_cast<uint32_t>(value));
} }
static inline void encode_int64(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id, int64_t value, [[nodiscard]] static inline uint8_t *encode_int32(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
bool force = false) { uint32_t field_id, int32_t value) {
encode_uint64(pos PROTO_ENCODE_DEBUG_ARG, field_id, static_cast<uint64_t>(value), force); if (value == 0)
return pos;
return encode_int32_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, value);
} }
static inline void encode_sint32(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id, [[nodiscard]] static inline uint8_t *encode_int64(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
int32_t value, bool force = false) { uint32_t field_id, int64_t value) {
encode_uint32(pos PROTO_ENCODE_DEBUG_ARG, field_id, encode_zigzag32(value), force); return encode_uint64(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, [[nodiscard]] static inline uint8_t *encode_int64_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
int64_t value, bool force = false) { uint32_t field_id, int64_t value) {
encode_uint64(pos PROTO_ENCODE_DEBUG_ARG, field_id, encode_zigzag64(value), force); return encode_uint64_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, static_cast<uint64_t>(value));
} }
/// Sub-message encoding: sync pos to buffer, delegate, get pos from return value. [[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.
template<typename T> template<typename T>
static inline void encode_sub_message(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, ProtoWriteBuffer &buffer, [[nodiscard]] static inline uint8_t *encode_sub_message(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t field_id, const T &value) { ProtoWriteBuffer &buffer, uint32_t field_id, const T &value) {
buffer.set_pos(pos); buffer.set_pos(pos);
buffer.encode_sub_message(field_id, value); buffer.encode_sub_message(field_id, value);
pos = buffer.get_pos(); return buffer.get_pos();
} }
template<typename T> template<typename T>
static inline void encode_optional_sub_message(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, [[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) { ProtoWriteBuffer &buffer, uint32_t field_id,
const T &value) {
buffer.set_pos(pos); buffer.set_pos(pos);
buffer.encode_optional_sub_message(field_id, value); buffer.encode_optional_sub_message(field_id, value);
pos = buffer.get_pos(); 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);
}
} }
}; };
#undef PROTO_OUTLINE_FOR_SIZE
#undef PROTO_FIXED32_BYTE_STORES
#ifdef HAS_PROTO_MESSAGE_DUMP #ifdef HAS_PROTO_MESSAGE_DUMP
/** /**
@@ -624,11 +692,12 @@ class DumpBuffer {
class ProtoMessage { class ProtoMessage {
public: public:
// Non-virtual defaults for messages with no fields. // Non-virtual defaults for messages with no fields; generated classes hide all four. The
// Concrete message classes hide these with their own implementations. // static encode_msg/calc_size_msg take const void * so &T::encode_msg needs no thunk.
// All call sites use templates to preserve the concrete type, so virtual static uint8_t *encode_msg(const void *self, ProtoWriteBuffer &buffer PROTO_ENCODE_DEBUG_PARAM) {
// dispatch is not needed. This eliminates per-message vtable entries for return buffer.get_pos();
// encode/calculate_size, saving ~1.3 KB of flash across all message types. }
static uint32_t calc_size_msg(const void *self) { return 0; }
uint8_t *encode(ProtoWriteBuffer &buffer PROTO_ENCODE_DEBUG_PARAM) const { return buffer.get_pos(); } uint8_t *encode(ProtoWriteBuffer &buffer PROTO_ENCODE_DEBUG_PARAM) const { return buffer.get_pos(); }
uint32_t calculate_size() const { return 0; } uint32_t calculate_size() const { return 0; }
#ifdef HAS_PROTO_MESSAGE_DUMP #ifdef HAS_PROTO_MESSAGE_DUMP
@@ -876,19 +945,14 @@ class ProtoSize {
// Implementation of methods that depend on ProtoSize being fully defined // 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. // 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) { template<typename T> inline void ProtoWriteBuffer::encode_sub_message(uint32_t field_id, const T &value) {
this->encode_sub_message(field_id, &value, &proto_encode_msg<T>); this->encode_sub_message(field_id, &value, &T::encode_msg);
} }
// Thin template wrapper; delegates to non-template core. // 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) { template<typename T> inline void ProtoWriteBuffer::encode_optional_sub_message(uint32_t field_id, const T &value) {
this->encode_optional_sub_message(field_id, value.calculate_size(), &value, &proto_encode_msg<T>); this->encode_optional_sub_message(field_id, T::calc_size_msg(&value), &value, &T::encode_msg);
} }
// Template decode_to_message - preserves concrete type so decode() resolves statically // Template decode_to_message - preserves concrete type so decode() resolves statically
-50
View File
@@ -389,7 +389,6 @@ class Application {
friend Component; friend Component;
friend class Scheduler; friend class Scheduler;
friend class LoopBlockingGuard; friend class LoopBlockingGuard;
friend class UnavoidableBlockingScope;
#ifdef USE_RUNTIME_STATS #ifdef USE_RUNTIME_STATS
friend class runtime_stats::RuntimeStatsCollector; friend class runtime_stats::RuntimeStatsCollector;
#endif #endif
@@ -632,55 +631,6 @@ class LoopBlockingGuard {
static void __attribute__((noinline, cold)) warn_blocking(uint32_t blocking_time); static void __attribute__((noinline, cold)) warn_blocking(uint32_t blocking_time);
}; };
/// Leaves a stretch of the current loop pass out of the blocking warning.
///
/// Only for work done from a loop pass that cannot be made shorter and
/// cannot be split across passes: turning on a radio, the first Wi-Fi
/// connect, a key generation whose cost is the algorithm itself. The warning
/// then keeps reporting everything else in the pass, and the component's
/// threshold does not ratchet up over the one step nothing can be done about.
///
/// Never use it to paper over a problem that can be solved. A slow driver
/// call, a loop that could be a state machine, a computation that could be
/// cached or deferred, a blocking read that could be polled: those are what
/// the warning exists to find, and wrapping them in this scope hides the
/// bug instead of fixing it. If in doubt, leave the warning in.
///
/// Only work timed by a LoopBlockingGuard is affected, that is a component's
/// loop() or a scheduler callback; setup() is not timed by the guard, so the
/// scope has no effect on the warning there. Main loop task only. The watchdog is not fed inside the
/// scope, so the work must finish within the watchdog timeout, or be paired
/// with a watchdog::WatchdogManager that raises the timeout for the same
/// stretch. Scopes may nest; the outermost one decides how much of the pass
/// is left out.
/// App.get_loop_component_start_time() reads later in the same pass return
/// the moved start, so elapsed time across the scope needs millis().
///
/// void MyComponent::loop() {
/// if (this->needs_key_) {
/// UnavoidableBlockingScope scope;
/// this->generate_key_();
/// }
/// }
class UnavoidableBlockingScope {
public:
UnavoidableBlockingScope() : started_(MillisInternal::get()), pass_start_(App.get_loop_component_start_time()) {}
~UnavoidableBlockingScope() {
// Move the pass start seen at entry forward by the time spent here, so an
// outer scope overrides an inner one instead of adding to it; never past
// now, which would underflow the guard's subtraction
const uint32_t now = MillisInternal::get();
const uint32_t moved = this->pass_start_ + (now - this->started_);
App.set_loop_component_start_time_(static_cast<int32_t>(now - moved) < 0 ? now : moved);
}
UnavoidableBlockingScope(const UnavoidableBlockingScope &) = delete;
UnavoidableBlockingScope &operator=(const UnavoidableBlockingScope &) = delete;
private:
uint32_t started_;
uint32_t pass_start_;
};
// Phase A: drain wake notifications and run the scheduler. Invoked on every // Phase A: drain wake notifications and run the scheduler. Invoked on every
// Application::loop() tick regardless of whether a component phase runs, so // Application::loop() tick regardless of whether a component phase runs, so
// scheduler items fire at their requested cadence even when the caller has // scheduler items fire at their requested cadence even when the caller has
-1
View File
@@ -51,7 +51,6 @@ class MillisInternal {
} }
friend class Application; friend class Application;
friend class LoopBlockingGuard; friend class LoopBlockingGuard;
friend class UnavoidableBlockingScope;
}; };
} // namespace esphome } // namespace esphome
+139 -74
View File
@@ -131,6 +131,12 @@ def force_str(force: bool) -> str:
return str(force).lower() return str(force).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))});"
class TypeInfo(ABC): class TypeInfo(ABC):
"""Base class for all type information.""" """Base class for all type information."""
@@ -264,14 +270,16 @@ class TypeInfo(ABC):
# write_raw_byte(tag) + raw encode instead of the full encode_* method, # write_raw_byte(tag) + raw encode instead of the full encode_* method,
# eliminating the zero-check branch and encode_field_raw indirection. # eliminating the zero-check branch and encode_field_raw indirection.
# {value} is replaced with the actual field expression. # {value} is replaced with the actual field expression.
RAW_ENCODE_MAP: dict[str, str] = { RAW_ENCODE_MAP: dict[str, tuple[str, str]] = {
"encode_uint32": "ProtoEncode::encode_varint_raw(pos, {value});", "encode_uint32": ("encode_varint_raw", "{value}"),
"encode_uint64": "ProtoEncode::encode_varint_raw_64(pos, {value});", "encode_uint64": ("encode_varint_raw_64", "{value}"),
"encode_sint32": "ProtoEncode::encode_varint_raw_short(pos, encode_zigzag32({value}));", "encode_sint32": ("encode_varint_raw_short", "encode_zigzag32({value})"),
"encode_sint64": "ProtoEncode::encode_varint_raw_64(pos, encode_zigzag64({value}));", "encode_sint64": ("encode_varint_raw_64", "encode_zigzag64({value})"),
"encode_int64": "ProtoEncode::encode_varint_raw_64(pos, static_cast<uint64_t>({value}));", "encode_int64": ("encode_varint_raw_64", "static_cast<uint64_t>({value})"),
"encode_bool": "ProtoEncode::write_raw_byte(pos, {value} ? 0x01 : 0x00);", "encode_bool": ("write_raw_byte", "{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: def _encode_with_precomputed_tag(self, value_expr: str) -> str | None:
"""Try to emit a precomputed-tag encode for a field. """Try to emit a precomputed-tag encode for a field.
@@ -288,12 +296,17 @@ class TypeInfo(ABC):
return None return None
max_val = self.max_value max_val = self.max_value
# Only use RAW_ENCODE_MAP for forced fields or fields with max_value # Only use RAW_ENCODE_MAP for forced fields or fields with max_value
raw_expr = None raw = None
if self.force or max_val is not None: if self.force or max_val is not None:
raw_expr = self.RAW_ENCODE_MAP.get(self.encode_func) raw = self.RAW_ENCODE_MAP.get(self.encode_func)
if raw_expr is None: if raw is None:
return None return None
body = f"ProtoEncode::write_raw_byte(pos, {tag});\n{raw_expr.format(value=value_expr)}" func, arg = raw
body = (
_encode_call("write_raw_byte", str(tag))
+ "\n"
+ _encode_call(func, arg.format(value=value_expr))
)
if self.force: if self.force:
return body return body
# Non-forced with max_value: inline zero-check + raw encode # Non-forced with max_value: inline zero-check + raw encode
@@ -314,23 +327,43 @@ class TypeInfo(ABC):
return None return None
# When max_len < 128, length varint is always 1 byte # When max_len < 128, length varint is always 1 byte
len_encode = ( len_encode = (
f"ProtoEncode::write_raw_byte(pos, static_cast<uint8_t>({len_expr}));" _encode_call("write_raw_byte", f"static_cast<uint8_t>({len_expr})")
if max_len is not None and max_len < 128 if max_len is not None and max_len < 128
else f"ProtoEncode::encode_varint_raw(pos, {len_expr});" else _encode_call("encode_varint_raw", 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_expr: str) -> str | None:
"""Single-byte tag fixed32 write, or None for multi-byte tags."""
tag = self.calculate_tag()
if tag >= 128:
return None
if self.force:
return _encode_call("write_tag_and_fixed32", str(tag), value_expr)
return ( return (
f"ProtoEncode::write_raw_byte(pos, {tag});\n" f"if (uint32_t raw = {value_expr}; raw != 0) [[likely]] {{\n"
f"{len_encode}\n" f" {_encode_call('write_tag_and_fixed32', str(tag), 'raw')}\n"
f"ProtoEncode::encode_raw(pos, {data_expr}, {len_expr});" "}"
) )
@property @property
def encode_content(self) -> str: def encode_content(self) -> str:
if result := self._encode_with_precomputed_tag(f"this->{self.field_name}"): value = f"this->{self.field_name}"
if result := self._encode_with_precomputed_tag(value):
return result return result
if self.force: if self.fixed32_value_template is not None and (
return f"ProtoEncode::{self.encode_func}(pos, {self.number}, this->{self.field_name}, true);" result := self._encode_fixed32_with_precomputed_tag(
return f"ProtoEncode::{self.encode_func}(pos, {self.number}, this->{self.field_name});" self.fixed32_value_template.format(value=value)
)
):
return result
return _encode_call(self.encode_func, str(self.number), value, force=self.force)
encode_func = None encode_func = None
@@ -635,6 +668,8 @@ class FloatType(FixedSizeTypeMixin, TypeInfo):
encode_func = "encode_float" encode_func = "encode_float"
wire_type = WireType.FIXED32 # Uses wire type 5 wire_type = WireType.FIXED32 # Uses wire type 5
fixed32_value_template = "float_to_raw({value})"
def dump(self, name: str) -> str: def dump(self, name: str) -> str:
o = f'snprintf(buffer, sizeof(buffer), "%g", {name});\n' o = f'snprintf(buffer, sizeof(buffer), "%g", {name});\n'
o += "out.append(buffer);" o += "out.append(buffer);"
@@ -697,11 +732,11 @@ class UInt64Type(VarintTypeMixin, TypeInfo):
return self._get_simple_size_calculation(name, force, "uint64") return self._get_simple_size_calculation(name, force, "uint64")
@property @property
def RAW_ENCODE_MAP(self) -> dict[str, str]: # noqa: N802 def RAW_ENCODE_MAP(self) -> dict[str, tuple[str, str]]: # noqa: N802
if self.mac_address: if self.mac_address:
return { return {
**TypeInfo.RAW_ENCODE_MAP, **TypeInfo.RAW_ENCODE_MAP,
"encode_uint64": "ProtoEncode::encode_varint_raw_48bit(pos, {value});", "encode_uint64": ("encode_varint_raw_48bit", "{value}"),
} }
return TypeInfo.RAW_ENCODE_MAP return TypeInfo.RAW_ENCODE_MAP
@@ -769,15 +804,7 @@ class Fixed32Type(FixedSizeTypeMixin, TypeInfo):
o += "out.append(buffer);" o += "out.append(buffer);"
return o return o
@property fixed32_value_template = "{value}"
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: def get_size_calculation(self, name: str, force: bool = False) -> str:
field_id_size = self.calculate_field_id_size() field_id_size = self.calculate_field_id_size()
@@ -851,9 +878,12 @@ class StringType(TypeInfo):
f"this->{self.field_name}_ref_.size()", f"this->{self.field_name}_ref_.size()",
): ):
return result return result
if self.force: return _encode_call(
return f"ProtoEncode::encode_string(pos, {self.number}, this->{self.field_name}_ref_, true);" "encode_string",
return f"ProtoEncode::encode_string(pos, {self.number}, this->{self.field_name}_ref_);" str(self.number),
f"this->{self.field_name}_ref_",
force=self.force,
)
def dump(self, name): def dump(self, name):
# If name is 'it', this is a repeated field element - always use string # If name is 'it', this is a repeated field element - always use string
@@ -951,7 +981,9 @@ class MessageType(TypeInfo):
@property @property
def encode_content(self) -> str: def encode_content(self) -> str:
# Sub-message encoding needs buffer for backpatch/sync # Sub-message encoding needs buffer for backpatch/sync
return f"ProtoEncode::{self.encode_func}(pos, buffer, {self.number}, this->{self.field_name});" return _encode_call(
self.encode_func, "buffer", str(self.number), f"this->{self.field_name}"
)
@property @property
def decode_length(self) -> str: def decode_length(self) -> str:
@@ -1058,9 +1090,13 @@ class BytesType(TypeInfo):
f"this->{self.field_name}_ptr_", f"this->{self.field_name}_len_" f"this->{self.field_name}_ptr_", f"this->{self.field_name}_len_"
): ):
return result return result
if self.force: return _encode_call(
return f"ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}_ptr_, this->{self.field_name}_len_, true);" "encode_bytes",
return f"ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}_ptr_, this->{self.field_name}_len_);" str(self.number),
f"this->{self.field_name}_ptr_",
f"this->{self.field_name}_len_",
force=self.force,
)
def dump(self, name: str) -> str: def dump(self, name: str) -> str:
ptr_dump = f"format_hex_pretty(this->{self.field_name}_ptr_, this->{self.field_name}_len_)" ptr_dump = f"format_hex_pretty(this->{self.field_name}_ptr_, this->{self.field_name}_len_)"
@@ -1170,9 +1206,13 @@ class PointerToBytesBufferType(PointerToBufferTypeBase):
f"this->{self.field_name}", f"this->{self.field_name}_len" f"this->{self.field_name}", f"this->{self.field_name}_len"
): ):
return result return result
if self.force: return _encode_call(
return f"ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}, this->{self.field_name}_len, true);" "encode_bytes",
return f"ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}, this->{self.field_name}_len);" str(self.number),
f"this->{self.field_name}",
f"this->{self.field_name}_len",
force=self.force,
)
@property @property
def decode_length_content(self) -> str | None: def decode_length_content(self) -> str | None:
@@ -1224,16 +1264,19 @@ class PointerToStringBufferType(PointerToBufferTypeBase):
if max_len is not None and max_len < 128 and self.force: if max_len is not None and max_len < 128 and self.force:
tag = self.calculate_tag() tag = self.calculate_tag()
if tag < 128: if tag < 128:
return f"ProtoEncode::encode_short_string_force(pos, {tag}, this->{self.field_name});" return _encode_call(
"encode_short_string_force", str(tag), f"this->{self.field_name}"
)
if result := self._encode_bytes_with_precomputed_tag( if result := self._encode_bytes_with_precomputed_tag(
f"this->{self.field_name}.c_str()", f"this->{self.field_name}.c_str()",
f"this->{self.field_name}.size()", f"this->{self.field_name}.size()",
): ):
return result return result
if self.force: return _encode_call(
return f"ProtoEncode::encode_string(pos, {self.number}, this->{self.field_name}, true);" "encode_string",
return ( str(self.number),
f"ProtoEncode::encode_string(pos, {self.number}, this->{self.field_name});" f"this->{self.field_name}",
force=self.force,
) )
@property @property
@@ -1421,9 +1464,13 @@ class FixedArrayBytesType(TypeInfo):
f"this->{self.field_name}", f"this->{self.field_name}_len", max_len=max_len f"this->{self.field_name}", f"this->{self.field_name}_len", max_len=max_len
): ):
return result return result
if self.force: return _encode_call(
return f"ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}, this->{self.field_name}_len, true);" "encode_bytes",
return f"ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}, this->{self.field_name}_len);" str(self.number),
f"this->{self.field_name}",
f"this->{self.field_name}_len",
force=self.force,
)
def dump(self, name: str) -> str: def dump(self, name: str) -> str:
return f"out.append(format_hex_pretty({name}, {name}_len));" return f"out.append(format_hex_pretty({name}, {name}_len));"
@@ -1520,9 +1567,9 @@ class EnumType(VarintTypeMixin, TypeInfo):
@property @property
def encode_content(self) -> str: def encode_content(self) -> str:
value_expr = f"static_cast<uint32_t>(this->{self.field_name})" value_expr = f"static_cast<uint32_t>(this->{self.field_name})"
if self.force: return _encode_call(
return f"ProtoEncode::{self.encode_func}(pos, {self.number}, {value_expr}, true);" self.encode_func, str(self.number), value_expr, force=self.force
return f"ProtoEncode::{self.encode_func}(pos, {self.number}, {value_expr});" )
def dump(self, name: str) -> str: def dump(self, name: str) -> str:
return f"out.append_p(proto_enum_to_string<{self.cpp_type}>({name}));" return f"out.append_p(proto_enum_to_string<{self.cpp_type}>({name}));"
@@ -1701,9 +1748,9 @@ def _generate_inline_encode_block(
lines = [] lines = []
lines.append(f"auto &sub_msg = {element};") lines.append(f"auto &sub_msg = {element};")
lines.append(f"ProtoEncode::write_raw_byte(pos, {tag});") lines.append(_encode_call("write_raw_byte", str(tag)))
lines.append("uint8_t *len_pos = pos;") lines.append("uint8_t *len_pos = pos;")
lines.append("ProtoEncode::reserve_byte(pos);") lines.append(_encode_call("reserve_byte"))
# Generate inline field encoding for each sub-message field # Generate inline field encoding for each sub-message field
for field in sub_desc.field: for field in sub_desc.field:
@@ -1775,17 +1822,22 @@ class FixedArrayRepeatedType(TypeInfo):
def _encode_element(self, element: str) -> str: def _encode_element(self, element: str) -> str:
"""Helper to generate encode statement for a single element.""" """Helper to generate encode statement for a single element."""
if isinstance(self._ti, EnumType): if isinstance(self._ti, EnumType):
return f"ProtoEncode::{self._ti.encode_func}(pos, {self.number}, static_cast<uint32_t>({element}), true);" return _encode_call(
self._ti.encode_func,
str(self.number),
f"static_cast<uint32_t>({element})",
force=True,
)
# Repeated message elements use encode_sub_message (force=true is default) # Repeated message elements use encode_sub_message (force=true is default)
if isinstance(self._ti, MessageType): if isinstance(self._ti, MessageType):
if _is_inline_encode(self._ti.cpp_type): if _is_inline_encode(self._ti.cpp_type):
return _generate_inline_encode_block( return _generate_inline_encode_block(
self.number, self._ti.cpp_type, element self.number, self._ti.cpp_type, element
) )
return f"ProtoEncode::encode_sub_message(pos, buffer, {self.number}, {element});" return _encode_call(
return ( "encode_sub_message", "buffer", str(self.number), element
f"ProtoEncode::{self._ti.encode_func}(pos, {self.number}, {element}, true);" )
) return _encode_call(self._ti.encode_func, str(self.number), element, force=True)
@property @property
def cpp_type(self) -> str: def cpp_type(self) -> str:
@@ -2137,13 +2189,18 @@ class RepeatedTypeInfo(TypeInfo):
def _encode_element_call(self, element: str) -> str: def _encode_element_call(self, element: str) -> str:
"""Helper to generate encode call for a single element.""" """Helper to generate encode call for a single element."""
if isinstance(self._ti, EnumType): if isinstance(self._ti, EnumType):
return f"ProtoEncode::{self._ti.encode_func}(pos, {self.number}, static_cast<uint32_t>({element}), true);" return _encode_call(
self._ti.encode_func,
str(self.number),
f"static_cast<uint32_t>({element})",
force=True,
)
# Repeated message elements use encode_sub_message (force=true is default) # Repeated message elements use encode_sub_message (force=true is default)
if isinstance(self._ti, MessageType): if isinstance(self._ti, MessageType):
return f"ProtoEncode::encode_sub_message(pos, buffer, {self.number}, {element});" return _encode_call(
return ( "encode_sub_message", "buffer", str(self.number), element
f"ProtoEncode::{self._ti.encode_func}(pos, {self.number}, {element}, true);" )
) return _encode_call(self._ti.encode_func, str(self.number), element, force=True)
@property @property
def encode_content(self) -> str: def encode_content(self) -> str:
@@ -2152,7 +2209,7 @@ class RepeatedTypeInfo(TypeInfo):
# Special handling for const char* elements (when container_no_template contains "const char") # Special handling for const char* elements (when container_no_template contains "const char")
if "const char" in self._container_no_template: if "const char" in self._container_no_template:
o = f"for (const char *it : *this->{self.field_name}) {{\n" o = f"for (const char *it : *this->{self.field_name}) {{\n"
o += f" ProtoEncode::{self._ti.encode_func}(pos, {self.number}, it, strlen(it), true);\n" o += f" {_encode_call(self._ti.encode_func, str(self.number), 'it', 'strlen(it)', force=True)}\n"
else: else:
o = f"for (const auto &it : *this->{self.field_name}) {{\n" o = f"for (const auto &it : *this->{self.field_name}) {{\n"
o += f" {self._encode_element_call('it')}\n" o += f" {self._encode_element_call('it')}\n"
@@ -2784,28 +2841,36 @@ def build_message_type(
) )
for line in encode for line in encode
] ]
o = f"{speed_attr}uint8_t *{desc.name}::encode(ProtoWriteBuffer &buffer PROTO_ENCODE_DEBUG_PARAM) const {{\n" 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 += " uint8_t *__restrict__ pos = buffer.get_pos();\n" o += " uint8_t *__restrict__ pos = buffer.get_pos();\n"
o += indent("\n".join(encode_debug)) + "\n" o += indent("\n".join(encode_debug)).replace("this->", "msg.") + "\n"
o += " return pos;\n" o += " return pos;\n"
o += "}\n" o += "}\n"
cpp += o cpp += o
prot = ( public_content.append(
"uint8_t *encode(ProtoWriteBuffer &buffer PROTO_ENCODE_DEBUG_PARAM) const;" "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"
"}"
) )
public_content.append(prot)
# If no fields to encode or message doesn't need encoding, the default implementation in ProtoMessage will be used # 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 # Add calculate_size method only if this message needs encoding and has fields
if needs_encode and size_calc and not is_inline_only: if needs_encode and size_calc and not is_inline_only:
o = f"{speed_attr}uint32_t {desc.name}::calculate_size() const {{\n" 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 += " uint32_t size = 0;\n" o += " uint32_t size = 0;\n"
o += indent("\n".join(size_calc)) + "\n" o += indent("\n".join(size_calc)).replace("this->", "msg.") + "\n"
o += " return size;\n" o += " return size;\n"
o += "}\n" o += "}\n"
cpp += o cpp += o
prot = "uint32_t calculate_size() const;" public_content.append("static uint32_t calc_size_msg(const void *self);")
public_content.append(prot) public_content.append(
"uint32_t calculate_size() const { return calc_size_msg(this); }"
)
# If no fields to calculate size for or message doesn't need encoding, the default implementation in ProtoMessage will be used # If no fields to calculate size for or message doesn't need encoding, the default implementation in ProtoMessage will be used
# dump_to method declaration in header # dump_to method declaration in header
@@ -59,7 +59,7 @@ static void verify_mac(uint64_t mac, size_t expected_bytes) {
#ifdef ESPHOME_DEBUG_API #ifdef ESPHOME_DEBUG_API
uint8_t *proto_debug_end_ = api_buf.data() + api_buf.size(); uint8_t *proto_debug_end_ = api_buf.data() + api_buf.size();
#endif #endif
ProtoEncode::encode_varint_raw_48bit(pos PROTO_ENCODE_DEBUG_ARG, mac); pos = ProtoEncode::encode_varint_raw_48bit(pos PROTO_ENCODE_DEBUG_ARG, mac);
size_t new_len = pos - api_buf.data(); size_t new_len = pos - api_buf.data();
EXPECT_EQ(new_len, expected_bytes) << "mac=0x" << std::hex << mac << std::dec; EXPECT_EQ(new_len, expected_bytes) << "mac=0x" << std::hex << mac << std::dec;
@@ -1,103 +0,0 @@
#include <gtest/gtest.h>
#include "esphome/core/application.h"
#include "esphome/core/hal.h"
namespace esphome {
// The scope must push the pass start forward by the time it covers and by
// nothing else, so the blocking guard sees only the work outside it
TEST(UnavoidableBlockingScope, ExcludesItsDurationFromThePass) {
const uint32_t pass_start = millis();
LoopBlockingGuard guard(nullptr, nullptr, pass_start);
ASSERT_EQ(App.get_loop_component_start_time(), pass_start);
const uint32_t before = millis();
{
UnavoidableBlockingScope scope;
delay(30);
}
const uint32_t excused = millis() - before;
const uint32_t moved = App.get_loop_component_start_time() - pass_start;
EXPECT_GE(moved, 30u);
EXPECT_LE(moved, excused);
}
TEST(UnavoidableBlockingScope, ZeroLengthScopeLeavesTheStartAlone) {
const uint32_t pass_start = millis();
LoopBlockingGuard guard(nullptr, nullptr, pass_start);
const uint32_t before = millis();
{ UnavoidableBlockingScope scope; }
EXPECT_LE(App.get_loop_component_start_time() - pass_start, millis() - before);
}
// Nested scopes leave out the outer span exactly once, and never move the
// start past now
TEST(UnavoidableBlockingScope, NestedScopesExcuseTheOuterSpanOnce) {
const uint32_t pass_start = millis();
LoopBlockingGuard guard(nullptr, nullptr, pass_start);
const uint32_t before = millis();
{
UnavoidableBlockingScope outer;
{
UnavoidableBlockingScope inner;
delay(30);
}
delay(5);
}
const uint32_t excused = millis() - before;
const uint32_t moved = App.get_loop_component_start_time() - pass_start;
EXPECT_GE(moved, 35u);
EXPECT_LE(moved, excused);
EXPECT_GE(static_cast<int32_t>(millis() - App.get_loop_component_start_time()), 0);
}
namespace {
// Static: the guard publishes the component to App and nothing clears it.
// One instance per test, since a ratcheted threshold is permanent
class DummyComponent : public Component {};
DummyComponent &blocking_test_component(size_t index) {
static DummyComponent components[2];
return components[index];
}
} // namespace
// The excused stretch must neither warn nor ratchet the component's threshold
TEST(UnavoidableBlockingScope, ExcusedStretchDoesNotRatchetTheThreshold) {
DummyComponent &component = blocking_test_component(0);
uint32_t threshold_before = 0;
component.should_warn_of_blocking(0, threshold_before);
{
LoopBlockingGuard guard(&component, nullptr, millis());
{
UnavoidableBlockingScope scope;
delay(WARN_IF_BLOCKING_OVER_CS * 10U + 20);
}
guard.finish();
}
uint32_t threshold_after = 0;
component.should_warn_of_blocking(0, threshold_after);
EXPECT_EQ(threshold_after, threshold_before);
}
// Work outside the scope is still measured and still ratchets
TEST(UnavoidableBlockingScope, WorkOutsideTheScopeStillRatchetsTheThreshold) {
DummyComponent &component = blocking_test_component(1);
uint32_t threshold_before = 0;
component.should_warn_of_blocking(0, threshold_before);
{
LoopBlockingGuard guard(&component, nullptr, millis());
{
UnavoidableBlockingScope scope;
delay(20);
}
delay(WARN_IF_BLOCKING_OVER_CS * 10U + 20);
guard.finish();
}
uint32_t threshold_after = 0;
component.should_warn_of_blocking(0, threshold_after);
EXPECT_GT(threshold_after, threshold_before);
}
} // namespace esphome
@@ -0,0 +1,11 @@
esphome:
name: api-empty-message-test
host:
api:
logger:
level: DEBUG
switch:
- platform: template
name: "Empty Message Switch"
optimistic: true
@@ -0,0 +1,58 @@
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');
+40
View File
@@ -57,6 +57,46 @@ async def wait_for_state(
return await asyncio.wait_for(future, timeout=timeout) 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]( def find_entity[T: EntityInfo](
entities: list[EntityInfo], entities: list[EntityInfo],
object_id_substring: str, object_id_substring: str,
@@ -0,0 +1,37 @@
"""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])
@@ -0,0 +1,78 @@
"""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",
),
)
@@ -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: def test_superseded_device_info_fields_still_encoded_and_sized() -> None:
"""Each superseded field must still be touched by DeviceInfoResponse's """Each superseded field must still be touched by DeviceInfoResponse's
generated encode() and calculate_size(), i.e. it is still put on the wire. generated encode_msg() and calc_size_msg(), i.e. it is still put on the wire.
""" """
encode_body = _extract_function_body(CPP_TEXT, "DeviceInfoResponse::encode") encode_body = _extract_function_body(CPP_TEXT, "DeviceInfoResponse::encode_msg")
size_body = _extract_function_body(CPP_TEXT, "DeviceInfoResponse::calculate_size") size_body = _extract_function_body(CPP_TEXT, "DeviceInfoResponse::calc_size_msg")
for field_name in SUPERSEDED_FIELDS: for field_name in SUPERSEDED_FIELDS:
assert f"this->{field_name}" in encode_body, ( assert f"msg.{field_name}" in encode_body, (
f"DeviceInfoResponse::encode() no longer references {field_name}. " f"DeviceInfoResponse::encode_msg() no longer references {field_name}. "
f"{DEPRECATED_FIELD_TRAP}" f"{DEPRECATED_FIELD_TRAP}"
) )
assert f"this->{field_name}" in size_body, ( assert f"msg.{field_name}" in size_body, (
f"DeviceInfoResponse::calculate_size() no longer references " f"DeviceInfoResponse::calc_size_msg() no longer references "
f"{field_name}. {DEPRECATED_FIELD_TRAP}" f"{field_name}. {DEPRECATED_FIELD_TRAP}"
) )
@@ -380,3 +380,13 @@ def test_api_version_minor_is_at_least_15() -> None:
"clients to see api_version >= 1.15 in HelloResponse before they will " "clients to see api_version >= 1.15 in HelloResponse before they will "
"ever request it." "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,9 +15,11 @@ import pytest
sys.path.insert(0, str(Path(__file__).parents[4] / "script" / "api_protobuf")) 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 from api_protobuf import ( # noqa: E402
MAX_MESSAGE_ID, MAX_MESSAGE_ID,
_make_ifdef_line, _make_ifdef_line,
create_field_type_info,
get_varint64_ifdef, get_varint64_ifdef,
validate_message_id, validate_message_id,
) )
@@ -43,7 +45,14 @@ UINT64 = descriptor_pb2.FieldDescriptorProto.TYPE_UINT64
INT64 = descriptor_pb2.FieldDescriptorProto.TYPE_INT64 INT64 = descriptor_pb2.FieldDescriptorProto.TYPE_INT64
SINT64 = descriptor_pb2.FieldDescriptorProto.TYPE_SINT64 SINT64 = descriptor_pb2.FieldDescriptorProto.TYPE_SINT64
UINT32 = descriptor_pb2.FieldDescriptorProto.TYPE_UINT32 UINT32 = descriptor_pb2.FieldDescriptorProto.TYPE_UINT32
INT32 = descriptor_pb2.FieldDescriptorProto.TYPE_INT32
SINT32 = descriptor_pb2.FieldDescriptorProto.TYPE_SINT32
FIXED64 = descriptor_pb2.FieldDescriptorProto.TYPE_FIXED64 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: def test_no_varint64_fields() -> None:
@@ -107,3 +116,69 @@ def test_message_id_at_maximum_is_accepted() -> None:
def test_message_id_above_maximum_is_rejected() -> None: def test_message_id_above_maximum_is_rejected() -> None:
with pytest.raises(ValueError, match="exceeds the plaintext"): with pytest.raises(ValueError, match="exceeds the plaintext"):
validate_message_id(MAX_MESSAGE_ID + 1, "TooBigMessage") 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