[api] Keep the protobuf write cursor in a register and emit sub-message tags inline (#20124)

This commit is contained in:
J. Nick Koston
2026-10-07 10:08:25 -10:00
committed by GitHub
parent 0323fd36d4
commit 3296bccd3c
12 changed files with 495 additions and 538 deletions
+1 -2
View File
@@ -2352,8 +2352,7 @@ bool APIConnection::send_message_(uint32_t payload_size, uint16_t message_type,
#endif
// Capacity reserved above, cannot fail
(void) shared_buf.resize(write_start + payload_size);
ProtoWriteBuffer buffer{&shared_buf, write_start};
uint8_t *end = encode_fn(msg, buffer PROTO_ENCODE_DEBUG_INIT(&shared_buf));
uint8_t *end = encode_fn(msg, shared_buf.data() + write_start PROTO_ENCODE_DEBUG_INIT(&shared_buf));
#ifdef ESPHOME_DEBUG_API
proto_check_encode_end(end, shared_buf.data() + shared_buf.size());
#else
+1 -1
View File
@@ -361,7 +361,7 @@ class APIConnection final : public APIServerConnectionBase {
void on_no_setup_connection();
// Function pointer type for type-erased message encoding
using MessageEncodeFn = uint8_t *(*) (const void *, ProtoWriteBuffer &PROTO_ENCODE_DEBUG_PARAM);
using MessageEncodeFn = ProtoEncodeFn;
// Function pointer type for type-erased size calculation
using CalculateSizeFn = uint32_t (*)(const void *);
@@ -45,8 +45,8 @@ inline uint16_t ESPHOME_ALWAYS_INLINE APIConnection::encode_to_buffer(uint32_t c
conn->fatal_out_of_memory_();
return 0;
}
ProtoWriteBuffer buffer{&shared_buf, shared_buf.size() - calculated_size};
uint8_t *end = encode_fn(msg, buffer PROTO_ENCODE_DEBUG_INIT(&shared_buf));
uint8_t *end =
encode_fn(msg, shared_buf.data() + shared_buf.size() - calculated_size 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
proto_check_encode_end(end, shared_buf.data() + shared_buf.size());
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+1 -2
View File
@@ -13,9 +13,8 @@ namespace esphome::api {
static const char *const TAG = "api.wizard";
uint8_t *wizard_encode_response(const void *self, ProtoWriteBuffer &buffer PROTO_ENCODE_DEBUG_PARAM) {
uint8_t *wizard_encode_response(const void *self, uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM) {
const auto &msg = *static_cast<const DeviceWizardResponse *>(self);
uint8_t *__restrict__ pos = buffer.get_pos();
if (msg.data_len == 0)
return pos;
pos = ProtoEncode::encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, 1, 2); // type 2: Length-delimited
+1 -1
View File
@@ -24,7 +24,7 @@ extern const uint8_t API_WIZARD_DATA[] PROGMEM;
/// Encodes a DeviceWizardResponse like the generated encoder would. The data is in flash, which ESP8266 can only read
/// with progmem_memcpy, so the generated encoder (a plain memcpy) cannot be used. Plain memcpy elsewhere.
uint8_t *wizard_encode_response(const void *self, ProtoWriteBuffer &buffer PROTO_ENCODE_DEBUG_PARAM);
uint8_t *wizard_encode_response(const void *self, uint8_t *pos PROTO_ENCODE_DEBUG_PARAM);
#ifdef USE_API_WIZARD_INPUTS
/// Where the entity id of an input is kept, found by the key the client uses for it.
+28 -42
View File
@@ -119,7 +119,7 @@ uint32_t ProtoDecodableMessage::count_repeated_field(const uint8_t *buffer, size
}
// Single-pass encode for repeated submessage elements (non-template core).
// Writes field tag, reserves 1 byte for length varint, encodes the submessage body,
// Reserves 1 byte for length varint, encodes the submessage body,
// then backpatches the actual length. For the common case (body < 128 bytes), this is
// just a single byte write with no memmove — all current repeated submessage types
// (BLE advertisements at ~47B, GATT descriptors at ~24B, service args, etc.) take
@@ -143,51 +143,34 @@ uint32_t ProtoDecodableMessage::count_repeated_field(const uint8_t *buffer, size
//
// After writing 2-byte varint at len_pos:
// [tag][v1][v2][body ..... body]
// ^-- pos_ = element end, within buffer
void ProtoWriteBuffer::encode_sub_message(uint32_t field_id, const void *value,
uint8_t *(*encode_fn)(const void *,
ProtoWriteBuffer &PROTO_ENCODE_DEBUG_PARAM)) {
this->encode_field_raw(field_id, 2);
// Reserve 1 byte for length varint (optimistic: submessage < 128 bytes)
uint8_t *len_pos = this->pos_;
this->debug_check_bounds_(1);
this->pos_++;
uint8_t *body_start = this->pos_;
this->pos_ = encode_fn(value, *this PROTO_ENCODE_DEBUG_INIT(this->buffer_));
uint32_t body_size = static_cast<uint32_t>(this->pos_ - body_start);
if (body_size < 128) [[likely]] {
// ^-- returned cursor = element end, within buffer
uint8_t *ProtoEncode::encode_sub_message_body(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, const void *value,
ProtoEncodeFn encode_fn) {
// Reserve 1 byte for the length varint (optimistic: submessage < 128 bytes)
uint8_t *len_pos = pos;
PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
uint8_t *body_start = pos + 1;
uint8_t *after_body = encode_fn(value, body_start PROTO_ENCODE_DEBUG_ARG);
uint32_t body_size = static_cast<uint32_t>(after_body - body_start);
if (body_size < VARINT_MAX_1_BYTE) [[likely]] {
// Common case: 1-byte varint, just backpatch
*len_pos = static_cast<uint8_t>(body_size);
return;
return after_body;
}
// Compute extra bytes needed for varint beyond the 1 already reserved
// Shift the body forward to make room for the extra length varint bytes
uint8_t extra = ProtoSize::varint(body_size) - 1;
// Shift body forward to make room for the extra varint bytes
this->debug_check_bounds_(extra);
PROTO_ENCODE_CHECK_BOUNDS(after_body, extra);
std::memmove(body_start + extra, body_start, body_size);
uint8_t *end = this->pos_ + extra;
// Write the full varint at len_pos
this->pos_ = len_pos;
this->encode_varint_raw(body_size);
this->pos_ = end;
(void) encode_varint_raw_loop(len_pos PROTO_ENCODE_DEBUG_ARG, body_size);
return after_body + extra;
}
// Non-template core for encode_optional_sub_message.
void ProtoWriteBuffer::encode_optional_sub_message(uint32_t field_id, uint32_t nested_size, const void *value,
uint8_t *(*encode_fn)(const void *,
ProtoWriteBuffer &PROTO_ENCODE_DEBUG_PARAM)) {
if (nested_size == 0)
return;
this->encode_field_raw(field_id, 2);
this->encode_varint_raw(nested_size);
#ifdef ESPHOME_DEBUG_API
uint8_t *start = this->pos_;
this->pos_ = encode_fn(value, *this PROTO_ENCODE_DEBUG_INIT(this->buffer_));
if (static_cast<uint32_t>(this->pos_ - start) != nested_size)
this->debug_check_encode_size_(field_id, nested_size, this->pos_ - start);
#else
this->pos_ = encode_fn(value, *this PROTO_ENCODE_DEBUG_INIT(this->buffer_));
#endif
uint8_t *ProtoEncode::encode_sized_sub_message_body(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t nested_size, const void *value, ProtoEncodeFn encode_fn) {
pos = encode_varint_raw(pos PROTO_ENCODE_DEBUG_ARG, nested_size);
return encode_fn(value, pos PROTO_ENCODE_DEBUG_ARG);
}
#ifdef ESPHOME_DEBUG_API
@@ -201,6 +184,14 @@ void proto_check_encode_end(const uint8_t *end, const uint8_t *expected) {
ESP_LOGE(TAG, "Proto encode ended %td bytes off the calculated size", end - expected);
abort();
}
void proto_check_sub_message_size(uint32_t field_id, uint32_t expected, const uint8_t *len_pos, const uint8_t *end) {
ptrdiff_t actual = end - (len_pos + ProtoSize::varint(expected));
if (actual == static_cast<ptrdiff_t>(expected))
return;
ESP_LOGE(TAG, "encode_message: size mismatch for field %" PRIu32 ": calculated=%" PRIu32 " actual=%td", field_id,
expected, actual);
abort();
}
void ProtoWriteBuffer::debug_check_bounds_(size_t bytes, const char *caller) {
if (this->pos_ + bytes > this->buffer_->data() + this->buffer_->size()) {
ESP_LOGE(TAG, "ProtoWriteBuffer bounds check failed in %s: bytes=%zu offset=%td buf_size=%zu", caller, bytes,
@@ -208,11 +199,6 @@ void ProtoWriteBuffer::debug_check_bounds_(size_t bytes, const char *caller) {
abort();
}
}
void ProtoWriteBuffer::debug_check_encode_size_(uint32_t field_id, uint32_t expected, ptrdiff_t actual) {
ESP_LOGE(TAG, "encode_message: size mismatch for field %" PRIu32 ": calculated=%" PRIu32 " actual=%td", field_id,
expected, actual);
abort();
}
#endif
+30 -41
View File
@@ -231,6 +231,8 @@ void proto_check_bounds_failed(const uint8_t *pos, size_t bytes, const uint8_t *
/// Aborts unless an encode body ended exactly where calculate_size() promised. A plain check rather than
/// assert(), so NDEBUG cannot switch it off.
void proto_check_encode_end(const uint8_t *end, const uint8_t *expected);
/// Aborts unless a sized sub-message (length prefix at len_pos) ended where its calculated size said.
void proto_check_sub_message_size(uint32_t field_id, uint32_t expected, const uint8_t *len_pos, const uint8_t *end);
#else
#define PROTO_ENCODE_DEBUG_PARAM
#define PROTO_ENCODE_DEBUG_ARG
@@ -263,21 +265,6 @@ class ProtoWriteBuffer {
* Following https://protobuf.dev/programming-guides/encoding/#structure
*/
void encode_field_raw(uint32_t field_id, uint32_t type) { this->encode_varint_raw(proto_tag(field_id, type)); }
/// Single-pass encode for repeated submessage elements.
/// Thin template wrapper; all buffer work is in the non-template core.
template<typename T> void encode_sub_message(uint32_t field_id, const T &value);
/// Encode an optional singular submessage field — skips if empty.
/// Thin template wrapper; all buffer work is in the non-template core.
template<typename T> void encode_optional_sub_message(uint32_t field_id, const T &value);
// NOLINTBEGIN(readability-identifier-naming)
// Non-template core for encode_sub_message — backpatch approach.
void encode_sub_message(uint32_t field_id, const void *value,
uint8_t *(*encode_fn)(const void *, ProtoWriteBuffer &PROTO_ENCODE_DEBUG_PARAM));
// Non-template core for encode_optional_sub_message.
void encode_optional_sub_message(uint32_t field_id, uint32_t nested_size, const void *value,
uint8_t *(*encode_fn)(const void *, ProtoWriteBuffer &PROTO_ENCODE_DEBUG_PARAM));
// NOLINTEND(readability-identifier-naming)
APIBuffer *get_buffer() const { return buffer_; }
uint8_t *get_pos() const { return pos_; }
void set_pos(uint8_t *pos) { pos_ = pos; }
@@ -288,7 +275,6 @@ class ProtoWriteBuffer {
#ifdef ESPHOME_DEBUG_API
void debug_check_bounds_(size_t bytes, const char *caller = __builtin_FUNCTION());
void debug_check_encode_size_(uint32_t field_id, uint32_t expected, ptrdiff_t actual);
#else
void debug_check_bounds_([[maybe_unused]] size_t bytes) {}
#endif
@@ -313,6 +299,9 @@ class ProtoWriteBuffer {
constexpr uint32_t VARINT_MAX_1_BYTE = 1 << 7; // 128
constexpr uint32_t VARINT_MAX_2_BYTE = 1 << 14; // 16384
/// Generated encode body: writes the fields at pos, returns the cursor past them.
using ProtoEncodeFn = uint8_t *(*) (const void *, uint8_t *PROTO_ENCODE_DEBUG_PARAM);
/// Static encode helpers for the generated encode bodies. Each takes the write cursor by value and
/// returns it advanced, so outlined calls at -Os chain through the return register instead of a
/// stack slot. Helpers without a _force suffix skip fields holding the proto3 default.
@@ -575,22 +564,36 @@ class ProtoEncode {
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.
/// Repeated sub-message element; the constant tag is written inline.
template<typename T>
[[nodiscard]] static inline uint8_t *encode_sub_message(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
ProtoWriteBuffer &buffer, uint32_t field_id, const T &value) {
buffer.set_pos(pos);
buffer.encode_sub_message(field_id, value);
return buffer.get_pos();
uint32_t field_id, const T &value) {
pos = encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 2);
return encode_sub_message_body(pos PROTO_ENCODE_DEBUG_ARG, &value, &T::encode_msg);
}
/// Singular sub-message field, skipped when it encodes to nothing.
template<typename T>
[[nodiscard]] static inline uint8_t *encode_optional_sub_message(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
ProtoWriteBuffer &buffer, uint32_t field_id,
const T &value) {
buffer.set_pos(pos);
buffer.encode_optional_sub_message(field_id, value);
return buffer.get_pos();
uint32_t field_id, const T &value) {
uint32_t nested_size = T::calc_size_msg(&value);
if (nested_size == 0)
return pos;
pos = encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 2);
#ifdef ESPHOME_DEBUG_API
uint8_t *end = encode_sized_sub_message_body(pos PROTO_ENCODE_DEBUG_ARG, nested_size, &value, &T::encode_msg);
proto_check_sub_message_size(field_id, nested_size, pos, end);
return end;
#else
return encode_sized_sub_message_body(pos, nested_size, &value, &T::encode_msg);
#endif
}
/// Length and body, length backpatched after the body is written.
[[nodiscard]] static uint8_t *encode_sub_message_body(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
const void *value, ProtoEncodeFn encode_fn);
/// Length and body for a precomputed size.
[[nodiscard]] static uint8_t *encode_sized_sub_message_body(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
uint32_t nested_size, const void *value,
ProtoEncodeFn encode_fn);
private:
/// Unaligned little endian store of four bytes: byte stores where the outlined helper lives (ESP-IDF, ARM
@@ -704,9 +707,7 @@ class ProtoMessage {
public:
// Non-virtual defaults for messages with no fields; generated classes hide all four. The
// static encode_msg/calc_size_msg take const void * so &T::encode_msg needs no thunk.
static uint8_t *encode_msg(const void *self, ProtoWriteBuffer &buffer PROTO_ENCODE_DEBUG_PARAM) {
return buffer.get_pos();
}
static uint8_t *encode_msg(const void *self, uint8_t *pos PROTO_ENCODE_DEBUG_PARAM) { return pos; }
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(); }
uint32_t calculate_size() const { return 0; }
@@ -960,18 +961,6 @@ class ProtoSize {
}
};
// Implementation of methods that depend on ProtoSize being fully defined
// Thin template wrapper; delegates to non-template core in proto.cpp.
template<typename T> inline void ProtoWriteBuffer::encode_sub_message(uint32_t field_id, const T &value) {
this->encode_sub_message(field_id, &value, &T::encode_msg);
}
// Thin template wrapper; delegates to non-template core.
template<typename T> inline void ProtoWriteBuffer::encode_optional_sub_message(uint32_t field_id, const T &value) {
this->encode_optional_sub_message(field_id, T::calc_size_msg(&value), &value, &T::encode_msg);
}
template<typename T> const char *proto_enum_to_string(T value);
// ProtoService removed — its methods were inlined into APIConnection.
+5 -7
View File
@@ -941,7 +941,7 @@ class MessageType(TypeInfo):
return False
def encode_element(self, number: int, element: str) -> str:
return _encode_call("encode_sub_message", "buffer", str(number), element)
return _encode_call("encode_sub_message", str(number), element)
@property
def cpp_type(self) -> str:
@@ -964,9 +964,8 @@ class MessageType(TypeInfo):
@property
def encode_content(self) -> str:
# Sub-message encoding needs buffer for backpatch/sync
return _encode_call(
self.encode_func, "buffer", str(self.number), f"this->{self.field_name}"
self.encode_func, str(self.number), f"this->{self.field_name}"
)
@property
@@ -2722,19 +2721,18 @@ def build_message_type(
)
for line in encode
]
o = f"{speed_attr}uint8_t *{desc.name}::encode_msg(const void *self, ProtoWriteBuffer &buffer PROTO_ENCODE_DEBUG_PARAM) {{\n"
o = f"{speed_attr}uint8_t *{desc.name}::encode_msg(const void *self, uint8_t *__restrict__ pos 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 += indent("\n".join(encode_debug)).replace("this->", "msg.") + "\n"
o += " return pos;\n"
o += "}\n"
cpp += o
public_content.append(
"static uint8_t *encode_msg(const void *self, ProtoWriteBuffer &buffer PROTO_ENCODE_DEBUG_PARAM);"
"static uint8_t *encode_msg(const void *self, uint8_t *pos 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"
" return encode_msg(this, buffer.get_pos() PROTO_ENCODE_DEBUG_ARG);\n"
"}"
)
# If no fields to encode or message doesn't need encoding, the default implementation in ProtoMessage will be used
@@ -0,0 +1,88 @@
#include <gtest/gtest.h>
#include <cstdint>
#include <vector>
#include "esphome/components/api/api_buffer.h"
#include "esphome/components/api/proto.h"
namespace esphome::api::testing {
// Sub-message whose body is raw bytes, so the body size is set directly by the test.
struct BlobMessage {
const uint8_t *data;
uint32_t len;
static uint8_t *encode_msg(const void *self, uint8_t *pos PROTO_ENCODE_DEBUG_PARAM) {
const auto &msg = *static_cast<const BlobMessage *>(self);
// An empty vector's data() may be null, and memcpy needs a valid source even for zero bytes
if (msg.len == 0)
return pos;
return ProtoEncode::encode_raw(pos PROTO_ENCODE_DEBUG_ARG, msg.data, msg.len);
}
static uint32_t calc_size_msg(const void *self) { return static_cast<const BlobMessage *>(self)->len; }
};
static void append_varint(std::vector<uint8_t> &out, uint32_t value) {
while (value > 0x7F) {
out.push_back(static_cast<uint8_t>(value | 0x80));
value >>= 7;
}
out.push_back(static_cast<uint8_t>(value));
}
static std::vector<uint8_t> make_body(uint32_t len) {
std::vector<uint8_t> body(len);
for (uint32_t i = 0; i < len; i++)
body[i] = static_cast<uint8_t>(i * 7 + 1);
return body;
}
static std::vector<uint8_t> expected_field(uint32_t field_id, const std::vector<uint8_t> &body) {
std::vector<uint8_t> out;
append_varint(out, (field_id << 3) | 2);
append_varint(out, body.size());
out.insert(out.end(), body.begin(), body.end());
return out;
}
// Encodes into a buffer of exactly the expected size, so any overrun trips ASan or the debug bounds check.
template<bool OPTIONAL> static void verify(uint32_t field_id, uint32_t body_len) {
std::vector<uint8_t> body = make_body(body_len);
std::vector<uint8_t> expected = expected_field(field_id, body);
if (OPTIONAL && body_len == 0)
expected.clear();
BlobMessage msg{body.data(), body_len};
APIBuffer buf;
ASSERT_TRUE(buf.resize(expected.empty() ? 1 : expected.size()));
uint8_t *pos = buf.data();
#ifdef ESPHOME_DEBUG_API
uint8_t *proto_debug_end_ = buf.data() + buf.size();
#endif
uint8_t *end;
if constexpr (OPTIONAL) {
end = ProtoEncode::encode_optional_sub_message(pos PROTO_ENCODE_DEBUG_ARG, field_id, msg);
} else {
end = ProtoEncode::encode_sub_message(pos PROTO_ENCODE_DEBUG_ARG, field_id, msg);
}
ASSERT_EQ(static_cast<size_t>(end - buf.data()), expected.size()) << "field " << field_id << " body " << body_len;
EXPECT_EQ(std::vector<uint8_t>(buf.data(), end), expected) << "field " << field_id << " body " << body_len;
}
TEST(ProtoSubMessage, OneByteTag) { verify<false>(4, 10); }
TEST(ProtoSubMessage, TwoByteTag) { verify<false>(20, 10); }
TEST(ProtoSubMessage, EmptyBody) { verify<false>(20, 0); }
TEST(ProtoSubMessage, LongestOneByteLength) { verify<false>(20, 127); }
// The length outgrows its reserved byte, so the body is moved forward
TEST(ProtoSubMessage, TwoByteLength) {
verify<false>(20, 128);
verify<false>(25, 200);
}
TEST(ProtoSubMessage, ThreeByteLength) { verify<false>(4, 20000); }
TEST(ProtoOptionalSubMessage, EmptyIsSkipped) { verify<true>(22, 0); }
TEST(ProtoOptionalSubMessage, OneByteTag) { verify<true>(1, 10); }
TEST(ProtoOptionalSubMessage, TwoByteTagAndLength) { verify<true>(22, 200); }
} // namespace esphome::api::testing
+2 -4
View File
@@ -47,16 +47,14 @@ const WizardInputEntry API_WIZARD_INPUTS[API_WIZARD_INPUT_COUNT] = {
using Bytes = std::vector<uint8_t>;
static Bytes encode(const ProtoMessage &msg, uint32_t (*calc)(const void *),
uint8_t *(*enc)(const void *, ProtoWriteBuffer &PROTO_ENCODE_DEBUG_PARAM)) {
static Bytes encode(const ProtoMessage &msg, uint32_t (*calc)(const void *), ProtoEncodeFn enc) {
APIBuffer buffer;
uint32_t size = calc(&msg);
EXPECT_TRUE(buffer.resize(size));
ProtoWriteBuffer writer(&buffer, 0);
#ifdef ESPHOME_DEBUG_API
uint8_t *proto_debug_end_ = buffer.data() + buffer.size();
#endif
uint8_t *end = enc(&msg, writer PROTO_ENCODE_DEBUG_ARG);
uint8_t *end = enc(&msg, buffer.data() PROTO_ENCODE_DEBUG_ARG);
EXPECT_EQ(static_cast<size_t>(end - buffer.data()), size);
return Bytes(buffer.data(), buffer.data() + size);
}