[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.
This commit is contained in:
J. Nick Koston
2026-09-07 09:21:54 +02:00
parent d34d3994e1
commit 3b14f4dfc8
5 changed files with 982 additions and 746 deletions
File diff suppressed because it is too large Load Diff
+185 -121
View File
@@ -287,19 +287,30 @@ class ProtoWriteBuffer {
uint8_t *pos_;
};
// Helpers that trade a call for flash on embedded targets. The host keeps them inline: there the
// fixed32 write is a single store, so outlining would only add a call (and CodSpeed counts it).
#ifdef USE_HOST
#define PROTO_OUTLINE_FOR_SIZE inline
#else
#define PROTO_OUTLINE_FOR_SIZE __attribute__((noinline))
#endif
// Varint encoding thresholds — used by both proto_encode_* free functions and ProtoSize.
constexpr uint32_t VARINT_MAX_1_BYTE = 1 << 7; // 128
constexpr uint32_t VARINT_MAX_2_BYTE = 1 << 14; // 16384
/// Static encode helpers for generated encode() functions.
/// Generated code hoists buffer.pos_ into a local uint8_t *__restrict__ pos,
/// then calls these methods which take pos by reference. No struct, no overhead.
/// For sub-messages, pos is synced back to buffer before the call and reloaded after.
/// Generated code hoists buffer.pos_ into a local uint8_t *__restrict__ pos and threads it through
/// these helpers by value: each one takes the cursor and returns the advanced cursor. At -Os the
/// compiler outlines most helpers, and returning the cursor lets consecutive calls chain through the
/// return register instead of spilling pos to a stack slot that a by-reference parameter would need.
/// Helpers without a _force suffix skip fields holding the proto3 default (zero or empty).
/// For sub-messages, pos is synced to the buffer before the call and read back after.
class ProtoEncode {
public:
/// Write a multi-byte varint directly through a pos pointer.
template<typename T>
static inline void encode_varint_raw_loop(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, T value) {
static inline uint8_t *encode_varint_raw_loop(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, T value) {
do {
PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
*pos++ = static_cast<uint8_t>(value | 0x80);
@@ -307,48 +318,49 @@ class ProtoEncode {
} while (value > 0x7F);
PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
*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,
uint32_t value) {
static inline uint8_t *ESPHOME_ALWAYS_INLINE encode_varint_raw(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t value) {
if (value < VARINT_MAX_1_BYTE) [[likely]] {
PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
*pos++ = static_cast<uint8_t>(value);
return;
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).
static inline void ESPHOME_ALWAYS_INLINE encode_varint_raw_short(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t value) {
static inline uint8_t *ESPHOME_ALWAYS_INLINE
encode_varint_raw_short(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint32_t value) {
if (value < VARINT_MAX_1_BYTE) [[likely]] {
PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
*pos++ = static_cast<uint8_t>(value);
return;
return pos;
}
if (value < VARINT_MAX_2_BYTE) [[likely]] {
PROTO_ENCODE_CHECK_BOUNDS(pos, 2);
*pos++ = static_cast<uint8_t>(value | 0x80);
*pos++ = static_cast<uint8_t>(value >> 7);
return;
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,
uint64_t value) {
static inline uint8_t *ESPHOME_ALWAYS_INLINE encode_varint_raw_64(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint64_t value) {
if (value < VARINT_MAX_1_BYTE) [[likely]] {
PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
*pos++ = static_cast<uint8_t>(value);
return;
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.
/// Real MAC addresses occupy the full 48 bits (OUI in upper 24), so the
/// fast path -- any non-zero bit in the top 6 of 48 -- emits exactly 7 bytes
/// with no per-byte branch. Falls back to the general loop otherwise.
/// Caller must guarantee value fits in 48 bits (checked in debug builds).
static inline void ESPHOME_ALWAYS_INLINE encode_varint_raw_48bit(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM,
uint64_t value) {
static inline uint8_t *ESPHOME_ALWAYS_INLINE
encode_varint_raw_48bit(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint64_t value) {
#ifdef ESPHOME_DEBUG_API
assert(value < (1ULL << (MAC_ADDRESS_SIZE * 8)) && "encode_varint_raw_48bit: value exceeds 48 bits");
#endif
@@ -363,38 +375,38 @@ class ProtoEncode {
pos[4] = static_cast<uint8_t>((value >> 28) | 0x80);
pos[5] = static_cast<uint8_t>((value >> 35) | 0x80);
pos[6] = static_cast<uint8_t>(value >> 42);
pos += 7;
return;
return pos + 7;
}
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,
uint32_t field_id, uint32_t type) {
encode_varint_raw(pos PROTO_ENCODE_DEBUG_ARG, (field_id << 3) | type);
static inline uint8_t *ESPHOME_ALWAYS_INLINE encode_field_raw(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t field_id, uint32_t type) {
return encode_varint_raw(pos PROTO_ENCODE_DEBUG_ARG, (field_id << 3) | type);
}
/// 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,
uint8_t b) {
static inline uint8_t *ESPHOME_ALWAYS_INLINE write_raw_byte(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint8_t b) {
PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
*pos++ = b;
return pos;
}
/// Reserve one byte for later backpatch (e.g., sub-message length).
/// Advances pos past the reserved byte without writing a value.
static inline void ESPHOME_ALWAYS_INLINE reserve_byte(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM) {
static inline uint8_t *ESPHOME_ALWAYS_INLINE reserve_byte(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM) {
PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
pos++;
return pos + 1;
}
/// 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,
const void *data, size_t len) {
static inline uint8_t *ESPHOME_ALWAYS_INLINE encode_raw(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
const void *data, size_t len) {
PROTO_ENCODE_CHECK_BOUNDS(pos, len);
std::memcpy(pos, data, len);
pos += len;
return pos + len;
}
/// Encode tag + 1-byte length + raw string data. For strings with max_data_length < 128.
/// Tag must be a single-byte varint (< 128). Always encodes (no zero check).
static inline void encode_short_string_force(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint8_t tag,
const StringRef &ref) {
static inline uint8_t *encode_short_string_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint8_t tag,
const StringRef &ref) {
#ifdef ESPHOME_DEBUG_API
assert(ref.size() < 128 && "encode_short_string_force: string exceeds max_data_length < 128");
#endif
@@ -402,135 +414,187 @@ class ProtoEncode {
pos[0] = tag;
pos[1] = static_cast<uint8_t>(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.
static inline void ESPHOME_ALWAYS_INLINE write_tag_and_fixed32(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM,
uint8_t tag, uint32_t value) {
/// Store a 32-bit value little-endian at an unaligned position. __builtin_memcpy stays a builtin even
/// under ESP-IDF's -fno-builtin-memcpy, so xtensa expands it to byte stores and the host to one store.
static inline void ESPHOME_ALWAYS_INLINE write_fixed32_le(uint8_t *__restrict__ pos, uint32_t value) {
#if __BYTE_ORDER__ == __ORDER_LITTLE_ENDIAN__
__builtin_memcpy(pos, &value, 4);
#else
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);
#endif
}
/// Write a precomputed tag byte + 32-bit little-endian value.
/// Outlined on embedded targets: one shared copy beats five inline stores at every fixed32/float field.
static PROTO_OUTLINE_FOR_SIZE uint8_t *write_tag_and_fixed32(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint8_t tag, uint32_t value) {
PROTO_ENCODE_CHECK_BOUNDS(pos, 5);
pos[0] = tag;
#if __BYTE_ORDER__ == __ORDER_LITTLE_ENDIAN__
std::memcpy(pos + 1, &value, 4);
#else
pos[1] = static_cast<uint8_t>(value & 0xFF);
pos[2] = static_cast<uint8_t>((value >> 8) & 0xFF);
pos[3] = static_cast<uint8_t>((value >> 16) & 0xFF);
pos[4] = static_cast<uint8_t>((value >> 24) & 0xFF);
#endif
pos += 5;
write_fixed32_le(pos + 1, value);
return pos + 5;
}
static inline void encode_string(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
const char *string, size_t len, bool force = false) {
if (len == 0 && !force)
return;
encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 2); // type 2: Length-delimited string
static inline uint8_t *encode_string_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
const char *string, size_t len) {
pos = encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 2); // type 2: Length-delimited string
// NOLINTNEXTLINE(readability-inconsistent-ifelse-braces) -- false positive on [[likely]] attribute
if (len < VARINT_MAX_1_BYTE) [[likely]] {
PROTO_ENCODE_CHECK_BOUNDS(pos, 1 + len);
*pos++ = static_cast<uint8_t>(len);
} else {
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);
}
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,
const std::string &value, bool force = false) {
encode_string(pos PROTO_ENCODE_DEBUG_ARG, field_id, value.data(), value.size(), force);
static inline uint8_t *encode_string(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
const char *string, size_t len) {
if (len == 0)
return pos;
return encode_string_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, string, len);
}
static inline void encode_string(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
const StringRef &ref, bool force = false) {
encode_string(pos PROTO_ENCODE_DEBUG_ARG, field_id, ref.c_str(), ref.size(), force);
static inline uint8_t *encode_string(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
const std::string &value) {
return encode_string(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,
const uint8_t *data, size_t len, bool force = false) {
encode_string(pos PROTO_ENCODE_DEBUG_ARG, field_id, reinterpret_cast<const char *>(data), len, force);
static inline uint8_t *encode_string_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
const std::string &value) {
return encode_string_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, value.data(), value.size());
}
static inline void encode_uint32(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
uint32_t value, bool force = false) {
if (value == 0 && !force)
return;
encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 0);
encode_varint_raw(pos PROTO_ENCODE_DEBUG_ARG, value);
static inline uint8_t *encode_string(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
const StringRef &ref) {
return encode_string(pos PROTO_ENCODE_DEBUG_ARG, field_id, ref.c_str(), ref.size());
}
static inline void encode_uint64(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
uint64_t value, bool force = false) {
if (value == 0 && !force)
return;
encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 0);
encode_varint_raw_64(pos PROTO_ENCODE_DEBUG_ARG, value);
static inline uint8_t *encode_string_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
const StringRef &ref) {
return encode_string_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, ref.c_str(), ref.size());
}
static inline void encode_bool(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id, bool value,
bool force = false) {
if (!value && !force)
return;
encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 0);
static inline uint8_t *encode_bytes(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
const uint8_t *data, size_t len) {
return encode_string(pos PROTO_ENCODE_DEBUG_ARG, field_id, reinterpret_cast<const char *>(data), len);
}
static inline uint8_t *encode_bytes_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
const uint8_t *data, size_t len) {
return encode_string_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, reinterpret_cast<const char *>(data), len);
}
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);
}
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);
}
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);
}
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);
}
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);
*pos++ = value ? 0x01 : 0x00;
return pos;
}
static inline void encode_fixed32(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
uint32_t value, bool force = false) {
if (value == 0 && !force)
return;
encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 5);
static inline uint8_t *encode_bool(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
bool value) {
if (!value)
return pos;
return encode_bool_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, value);
}
/// Generic tag + fixed32 writer for tags that need more than one byte; single-byte tags use
/// write_tag_and_fixed32. Outlined under -Os for the same reason.
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);
#if __BYTE_ORDER__ == __ORDER_LITTLE_ENDIAN__
std::memcpy(pos, &value, 4);
pos += 4;
#else
*pos++ = (value >> 0) & 0xFF;
*pos++ = (value >> 8) & 0xFF;
*pos++ = (value >> 16) & 0xFF;
*pos++ = (value >> 24) & 0xFF;
#endif
write_fixed32_le(pos, value);
return pos + 4;
}
static inline uint8_t *encode_fixed32(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
uint32_t value) {
if (value == 0)
return pos;
return encode_fixed32_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, value);
}
// NOTE: Wire type 1 (64-bit fixed: double, fixed64, sfixed64) is intentionally
// not supported to reduce overhead on embedded systems. All ESPHome devices are
// 32-bit microcontrollers where 64-bit operations are expensive. If 64-bit support
// is needed in the future, the necessary encoding/decoding functions must be added.
static inline void encode_float(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id, float value,
bool force = false) {
uint32_t raw = float_to_raw(value);
if (raw == 0 && !force)
return;
encode_fixed32(pos PROTO_ENCODE_DEBUG_ARG, field_id, raw);
static inline uint8_t *encode_float(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
float value) {
return encode_fixed32(pos PROTO_ENCODE_DEBUG_ARG, field_id, float_to_raw(value));
}
static inline void encode_int32(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id, int32_t value,
bool force = false) {
static inline uint8_t *encode_float_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
float value) {
return encode_fixed32_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, float_to_raw(value));
}
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) {
// negative int32 is always 10 byte long
encode_uint64(pos PROTO_ENCODE_DEBUG_ARG, field_id, static_cast<uint64_t>(value), force);
return;
return encode_uint64_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, static_cast<uint64_t>(value));
}
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,
bool force = false) {
encode_uint64(pos PROTO_ENCODE_DEBUG_ARG, field_id, static_cast<uint64_t>(value), force);
static inline uint8_t *encode_int32(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
int32_t value) {
if (value == 0)
return pos;
return encode_int32_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, value);
}
static inline void encode_sint32(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
int32_t value, bool force = false) {
encode_uint32(pos PROTO_ENCODE_DEBUG_ARG, field_id, encode_zigzag32(value), force);
static inline uint8_t *encode_int64(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
int64_t value) {
return encode_uint64(pos PROTO_ENCODE_DEBUG_ARG, field_id, static_cast<uint64_t>(value));
}
static inline void encode_sint64(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
int64_t value, bool force = false) {
encode_uint64(pos PROTO_ENCODE_DEBUG_ARG, field_id, encode_zigzag64(value), force);
static inline uint8_t *encode_int64_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
int64_t value) {
return encode_uint64_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, static_cast<uint64_t>(value));
}
/// Sub-message encoding: sync pos to buffer, delegate, get pos from return value.
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));
}
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));
}
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));
}
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>
static inline void encode_sub_message(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, ProtoWriteBuffer &buffer,
uint32_t field_id, const T &value) {
static inline uint8_t *encode_sub_message(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
ProtoWriteBuffer &buffer, uint32_t field_id, const T &value) {
buffer.set_pos(pos);
buffer.encode_sub_message(field_id, value);
pos = buffer.get_pos();
return buffer.get_pos();
}
template<typename T>
static inline void encode_optional_sub_message(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM,
ProtoWriteBuffer &buffer, uint32_t field_id, const T &value) {
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) {
buffer.set_pos(pos);
buffer.encode_optional_sub_message(field_id, value);
pos = buffer.get_pos();
return buffer.get_pos();
}
};
+63 -48
View File
@@ -265,12 +265,12 @@ class TypeInfo(ABC):
# eliminating the zero-check branch and encode_field_raw indirection.
# {value} is replaced with the actual field expression.
RAW_ENCODE_MAP: dict[str, str] = {
"encode_uint32": "ProtoEncode::encode_varint_raw(pos, {value});",
"encode_uint64": "ProtoEncode::encode_varint_raw_64(pos, {value});",
"encode_sint32": "ProtoEncode::encode_varint_raw_short(pos, encode_zigzag32({value}));",
"encode_sint64": "ProtoEncode::encode_varint_raw_64(pos, encode_zigzag64({value}));",
"encode_int64": "ProtoEncode::encode_varint_raw_64(pos, static_cast<uint64_t>({value}));",
"encode_bool": "ProtoEncode::write_raw_byte(pos, {value} ? 0x01 : 0x00);",
"encode_uint32": "pos = ProtoEncode::encode_varint_raw(pos, {value});",
"encode_uint64": "pos = ProtoEncode::encode_varint_raw_64(pos, {value});",
"encode_sint32": "pos = ProtoEncode::encode_varint_raw_short(pos, encode_zigzag32({value}));",
"encode_sint64": "pos = ProtoEncode::encode_varint_raw_64(pos, encode_zigzag64({value}));",
"encode_int64": "pos = ProtoEncode::encode_varint_raw_64(pos, static_cast<uint64_t>({value}));",
"encode_bool": "pos = ProtoEncode::write_raw_byte(pos, {value} ? 0x01 : 0x00);",
}
def _encode_with_precomputed_tag(self, value_expr: str) -> str | None:
@@ -293,7 +293,7 @@ class TypeInfo(ABC):
raw_expr = self.RAW_ENCODE_MAP.get(self.encode_func)
if raw_expr is None:
return None
body = f"ProtoEncode::write_raw_byte(pos, {tag});\n{raw_expr.format(value=value_expr)}"
body = f"pos = ProtoEncode::write_raw_byte(pos, {tag});\n{raw_expr.format(value=value_expr)}"
if self.force:
return body
# Non-forced with max_value: inline zero-check + raw encode
@@ -314,14 +314,14 @@ class TypeInfo(ABC):
return None
# When max_len < 128, length varint is always 1 byte
len_encode = (
f"ProtoEncode::write_raw_byte(pos, static_cast<uint8_t>({len_expr}));"
f"pos = ProtoEncode::write_raw_byte(pos, static_cast<uint8_t>({len_expr}));"
if max_len is not None and max_len < 128
else f"ProtoEncode::encode_varint_raw(pos, {len_expr});"
else f"pos = ProtoEncode::encode_varint_raw(pos, {len_expr});"
)
return (
f"ProtoEncode::write_raw_byte(pos, {tag});\n"
f"pos = ProtoEncode::write_raw_byte(pos, {tag});\n"
f"{len_encode}\n"
f"ProtoEncode::encode_raw(pos, {data_expr}, {len_expr});"
f"pos = ProtoEncode::encode_raw(pos, {data_expr}, {len_expr});"
)
@property
@@ -329,8 +329,8 @@ class TypeInfo(ABC):
if result := self._encode_with_precomputed_tag(f"this->{self.field_name}"):
return result
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});"
return f"pos = ProtoEncode::{self.encode_func}_force(pos, {self.number}, this->{self.field_name});"
return f"pos = ProtoEncode::{self.encode_func}(pos, {self.number}, this->{self.field_name});"
encode_func = None
@@ -635,6 +635,21 @@ class FloatType(FixedSizeTypeMixin, TypeInfo):
encode_func = "encode_float"
wire_type = WireType.FIXED32 # Uses wire type 5
@property
def encode_content(self) -> str:
tag = self.calculate_tag()
if tag >= 128:
return super().encode_content
# Single-byte tag: share the outlined tag+fixed32 writer instead of the generic helper
value = f"float_to_raw(this->{self.field_name})"
if self.force:
return f"pos = ProtoEncode::write_tag_and_fixed32(pos, {tag}, {value});"
return (
f"if (uint32_t raw = {value}; raw != 0) [[likely]] {{\n"
f" pos = ProtoEncode::write_tag_and_fixed32(pos, {tag}, raw);\n"
"}"
)
def dump(self, name: str) -> str:
o = f'snprintf(buffer, sizeof(buffer), "%g", {name});\n'
o += "out.append(buffer);"
@@ -701,7 +716,7 @@ class UInt64Type(VarintTypeMixin, TypeInfo):
if self.mac_address:
return {
**TypeInfo.RAW_ENCODE_MAP,
"encode_uint64": "ProtoEncode::encode_varint_raw_48bit(pos, {value});",
"encode_uint64": "pos = ProtoEncode::encode_varint_raw_48bit(pos, {value});",
}
return TypeInfo.RAW_ENCODE_MAP
@@ -772,12 +787,16 @@ class Fixed32Type(FixedSizeTypeMixin, TypeInfo):
@property
def encode_content(self) -> str:
tag = self.calculate_tag()
if self.force and tag < 128:
# Emit combined tag+value write: precomputed tag + direct memcpy
return f"ProtoEncode::write_tag_and_fixed32(pos, {tag}, this->{self.field_name});"
if tag >= 128:
return super().encode_content
# Single-byte tag: share the outlined tag+fixed32 writer instead of the generic helper
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});"
return f"pos = ProtoEncode::write_tag_and_fixed32(pos, {tag}, this->{self.field_name});"
return (
f"if (this->{self.field_name} != 0) [[likely]] {{\n"
f" pos = ProtoEncode::write_tag_and_fixed32(pos, {tag}, this->{self.field_name});\n"
"}"
)
def get_size_calculation(self, name: str, force: bool = False) -> str:
field_id_size = self.calculate_field_id_size()
@@ -852,8 +871,8 @@ class StringType(TypeInfo):
):
return result
if self.force:
return f"ProtoEncode::encode_string(pos, {self.number}, this->{self.field_name}_ref_, true);"
return f"ProtoEncode::encode_string(pos, {self.number}, this->{self.field_name}_ref_);"
return f"pos = ProtoEncode::encode_string_force(pos, {self.number}, this->{self.field_name}_ref_);"
return f"pos = ProtoEncode::encode_string(pos, {self.number}, this->{self.field_name}_ref_);"
def dump(self, name):
# If name is 'it', this is a repeated field element - always use string
@@ -951,7 +970,7 @@ class MessageType(TypeInfo):
@property
def encode_content(self) -> str:
# Sub-message encoding needs buffer for backpatch/sync
return f"ProtoEncode::{self.encode_func}(pos, buffer, {self.number}, this->{self.field_name});"
return f"pos = ProtoEncode::{self.encode_func}(pos, buffer, {self.number}, this->{self.field_name});"
@property
def decode_length(self) -> str:
@@ -1059,8 +1078,8 @@ class BytesType(TypeInfo):
):
return result
if self.force:
return f"ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}_ptr_, this->{self.field_name}_len_, true);"
return f"ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}_ptr_, this->{self.field_name}_len_);"
return f"pos = ProtoEncode::encode_bytes_force(pos, {self.number}, this->{self.field_name}_ptr_, this->{self.field_name}_len_);"
return f"pos = ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}_ptr_, this->{self.field_name}_len_);"
def dump(self, name: str) -> str:
ptr_dump = f"format_hex_pretty(this->{self.field_name}_ptr_, this->{self.field_name}_len_)"
@@ -1171,8 +1190,8 @@ class PointerToBytesBufferType(PointerToBufferTypeBase):
):
return result
if self.force:
return f"ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}, this->{self.field_name}_len, true);"
return f"ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}, this->{self.field_name}_len);"
return f"pos = ProtoEncode::encode_bytes_force(pos, {self.number}, this->{self.field_name}, this->{self.field_name}_len);"
return f"pos = ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}, this->{self.field_name}_len);"
@property
def decode_length_content(self) -> str | None:
@@ -1224,17 +1243,15 @@ class PointerToStringBufferType(PointerToBufferTypeBase):
if max_len is not None and max_len < 128 and self.force:
tag = self.calculate_tag()
if tag < 128:
return f"ProtoEncode::encode_short_string_force(pos, {tag}, this->{self.field_name});"
return f"pos = ProtoEncode::encode_short_string_force(pos, {tag}, this->{self.field_name});"
if result := self._encode_bytes_with_precomputed_tag(
f"this->{self.field_name}.c_str()",
f"this->{self.field_name}.size()",
):
return result
if self.force:
return f"ProtoEncode::encode_string(pos, {self.number}, this->{self.field_name}, true);"
return (
f"ProtoEncode::encode_string(pos, {self.number}, this->{self.field_name});"
)
return f"pos = ProtoEncode::encode_string_force(pos, {self.number}, this->{self.field_name});"
return f"pos = ProtoEncode::encode_string(pos, {self.number}, this->{self.field_name});"
@property
def decode_length_content(self) -> str | None:
@@ -1422,8 +1439,8 @@ class FixedArrayBytesType(TypeInfo):
):
return result
if self.force:
return f"ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}, this->{self.field_name}_len, true);"
return f"ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}, this->{self.field_name}_len);"
return f"pos = ProtoEncode::encode_bytes_force(pos, {self.number}, this->{self.field_name}, this->{self.field_name}_len);"
return f"pos = ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}, this->{self.field_name}_len);"
def dump(self, name: str) -> str:
return f"out.append(format_hex_pretty({name}, {name}_len));"
@@ -1521,8 +1538,10 @@ class EnumType(VarintTypeMixin, TypeInfo):
def encode_content(self) -> str:
value_expr = f"static_cast<uint32_t>(this->{self.field_name})"
if self.force:
return f"ProtoEncode::{self.encode_func}(pos, {self.number}, {value_expr}, true);"
return f"ProtoEncode::{self.encode_func}(pos, {self.number}, {value_expr});"
return f"pos = ProtoEncode::{self.encode_func}_force(pos, {self.number}, {value_expr});"
return (
f"pos = ProtoEncode::{self.encode_func}(pos, {self.number}, {value_expr});"
)
def dump(self, name: str) -> str:
return f"out.append_p(proto_enum_to_string<{self.cpp_type}>({name}));"
@@ -1701,9 +1720,9 @@ def _generate_inline_encode_block(
lines = []
lines.append(f"auto &sub_msg = {element};")
lines.append(f"ProtoEncode::write_raw_byte(pos, {tag});")
lines.append(f"pos = ProtoEncode::write_raw_byte(pos, {tag});")
lines.append("uint8_t *len_pos = pos;")
lines.append("ProtoEncode::reserve_byte(pos);")
lines.append("pos = ProtoEncode::reserve_byte(pos);")
# Generate inline field encoding for each sub-message field
for field in sub_desc.field:
@@ -1775,17 +1794,15 @@ class FixedArrayRepeatedType(TypeInfo):
def _encode_element(self, element: str) -> str:
"""Helper to generate encode statement for a single element."""
if isinstance(self._ti, EnumType):
return f"ProtoEncode::{self._ti.encode_func}(pos, {self.number}, static_cast<uint32_t>({element}), true);"
return f"pos = ProtoEncode::{self._ti.encode_func}_force(pos, {self.number}, static_cast<uint32_t>({element}));"
# Repeated message elements use encode_sub_message (force=true is default)
if isinstance(self._ti, MessageType):
if _is_inline_encode(self._ti.cpp_type):
return _generate_inline_encode_block(
self.number, self._ti.cpp_type, element
)
return f"ProtoEncode::encode_sub_message(pos, buffer, {self.number}, {element});"
return (
f"ProtoEncode::{self._ti.encode_func}(pos, {self.number}, {element}, true);"
)
return f"pos = ProtoEncode::encode_sub_message(pos, buffer, {self.number}, {element});"
return f"pos = ProtoEncode::{self._ti.encode_func}_force(pos, {self.number}, {element});"
@property
def cpp_type(self) -> str:
@@ -2137,13 +2154,11 @@ class RepeatedTypeInfo(TypeInfo):
def _encode_element_call(self, element: str) -> str:
"""Helper to generate encode call for a single element."""
if isinstance(self._ti, EnumType):
return f"ProtoEncode::{self._ti.encode_func}(pos, {self.number}, static_cast<uint32_t>({element}), true);"
return f"pos = ProtoEncode::{self._ti.encode_func}_force(pos, {self.number}, static_cast<uint32_t>({element}));"
# Repeated message elements use encode_sub_message (force=true is default)
if isinstance(self._ti, MessageType):
return f"ProtoEncode::encode_sub_message(pos, buffer, {self.number}, {element});"
return (
f"ProtoEncode::{self._ti.encode_func}(pos, {self.number}, {element}, true);"
)
return f"pos = ProtoEncode::encode_sub_message(pos, buffer, {self.number}, {element});"
return f"pos = ProtoEncode::{self._ti.encode_func}_force(pos, {self.number}, {element});"
@property
def encode_content(self) -> str:
@@ -2152,7 +2167,7 @@ class RepeatedTypeInfo(TypeInfo):
# Special handling for const char* elements (when container_no_template contains "const char")
if "const char" in self._container_no_template:
o = f"for (const char *it : *this->{self.field_name}) {{\n"
o += f" ProtoEncode::{self._ti.encode_func}(pos, {self.number}, it, strlen(it), true);\n"
o += f" pos = ProtoEncode::{self._ti.encode_func}_force(pos, {self.number}, it, strlen(it));\n"
else:
o = f"for (const auto &it : *this->{self.field_name}) {{\n"
o += f" {self._encode_element_call('it')}\n"
@@ -59,7 +59,7 @@ static void verify_mac(uint64_t mac, size_t expected_bytes) {
#ifdef ESPHOME_DEBUG_API
uint8_t *proto_debug_end_ = api_buf.data() + api_buf.size();
#endif
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();
EXPECT_EQ(new_len, expected_bytes) << "mac=0x" << std::hex << mac << std::dec;
@@ -15,9 +15,11 @@ import pytest
sys.path.insert(0, str(Path(__file__).parents[4] / "script" / "api_protobuf"))
import aioesphomeapi.api_options_pb2 as pb # noqa: E402
from api_protobuf import ( # noqa: E402
MAX_MESSAGE_ID,
_make_ifdef_line,
create_field_type_info,
get_varint64_ifdef,
validate_message_id,
)
@@ -107,3 +109,80 @@ def test_message_id_at_maximum_is_accepted() -> None:
def test_message_id_above_maximum_is_rejected() -> None:
with pytest.raises(ValueError, match="exceeds the plaintext"):
validate_message_id(MAX_MESSAGE_ID + 1, "TooBigMessage")
def _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 = 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
ti = create_field_type_info(field, needs_decode=False, needs_encode=True)
return ti.encode_content
SCALAR_TYPES = [
descriptor_pb2.FieldDescriptorProto.TYPE_BOOL,
descriptor_pb2.FieldDescriptorProto.TYPE_UINT32,
descriptor_pb2.FieldDescriptorProto.TYPE_INT32,
descriptor_pb2.FieldDescriptorProto.TYPE_UINT64,
descriptor_pb2.FieldDescriptorProto.TYPE_INT64,
descriptor_pb2.FieldDescriptorProto.TYPE_SINT32,
descriptor_pb2.FieldDescriptorProto.TYPE_FLOAT,
descriptor_pb2.FieldDescriptorProto.TYPE_FIXED32,
descriptor_pb2.FieldDescriptorProto.TYPE_STRING,
descriptor_pb2.FieldDescriptorProto.TYPE_BYTES,
]
@pytest.mark.parametrize("field_type", SCALAR_TYPES)
@pytest.mark.parametrize("force", [False, True])
@pytest.mark.parametrize("repeated", [False, True])
def test_encode_statements_assign_the_returned_cursor(
field_type: int, force: bool, repeated: bool
) -> None:
"""Every ProtoEncode call must take pos by value and store the returned cursor."""
content = _encode_field(field_type, force=force, repeated=repeated)
calls = [line.strip() for line in content.splitlines() if "ProtoEncode::" in line]
assert calls, content
for call in calls:
assert call.startswith("pos = ProtoEncode::"), call
assert ", true)" not in content, content
@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
FLOAT = descriptor_pb2.FieldDescriptorProto.TYPE_FLOAT
FIXED32 = descriptor_pb2.FieldDescriptorProto.TYPE_FIXED32
@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