mirror of
https://github.com/esphome/esphome.git
synced 2026-09-08 22:08:49 +00:00
Compare commits
19
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
455e1d6374 | ||
|
|
3af1d50bce | ||
|
|
842f354a05 | ||
|
|
af9b59d4bd | ||
|
|
55fc5a10de | ||
|
|
79927b918b | ||
|
|
1b070629bc | ||
|
|
c4e1360cdf | ||
|
|
7c774699d7 | ||
|
|
ea71a24a9b | ||
|
|
709a1e1eb6 | ||
|
|
822b701792 | ||
|
|
d2e4d2c46a | ||
|
|
adbbda4072 | ||
|
|
252bf6ea6a | ||
|
|
8ec9305688 | ||
|
|
490aca17e6 | ||
|
|
b77e2441d4 | ||
|
|
3b14f4dfc8 |
@@ -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)
|
||||||
|
|||||||
@@ -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;
|
||||||
}
|
}
|
||||||
|
|||||||
+1646
-1348
File diff suppressed because it is too large
Load Diff
+594
-198
File diff suppressed because it is too large
Load Diff
+198
-134
@@ -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
|
||||||
|
|||||||
@@ -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;
|
||||||
|
|||||||
@@ -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');
|
||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user