mirror of
https://github.com/esphome/esphome.git
synced 2026-09-11 07:17:33 +00:00
Compare commits
12
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
12795852ca | ||
|
|
fa00b58bbd | ||
|
|
b7f6eb0a7d | ||
|
|
d29d880045 | ||
|
|
f17a1023f2 | ||
|
|
48a92430db | ||
|
|
89c222eb94 | ||
|
|
be9b56c5e0 | ||
|
|
39cdcf6d6f | ||
|
|
afd718ec3e | ||
|
|
3bd2faf4a9 | ||
|
|
a772f17235 |
@@ -2269,12 +2269,7 @@ bool APIConnection::send_message_(uint32_t payload_size, uint16_t message_type,
|
||||
// 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));
|
||||
#ifdef ESPHOME_DEBUG_API
|
||||
assert(end == shared_buf.data() + shared_buf.size());
|
||||
#else
|
||||
(void) end;
|
||||
#endif
|
||||
encode_fn(msg, buffer PROTO_ENCODE_DEBUG_INIT(&shared_buf));
|
||||
return this->send_buffer(ProtoWriteBuffer{&shared_buf}, message_type);
|
||||
}
|
||||
// encode_to_buffer is defined inline in api_connection.h (ESPHOME_ALWAYS_INLINE)
|
||||
|
||||
@@ -346,7 +346,11 @@ class APIConnection final : public APIServerConnectionBase {
|
||||
/// 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.
|
||||
template<typename T> [[nodiscard]] bool send_message(const T &msg) {
|
||||
return this->send_message_(T::calc_size_msg(&msg), T::MESSAGE_TYPE, &T::encode_msg, &msg);
|
||||
if constexpr (T::ESTIMATED_SIZE == 0) {
|
||||
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.
|
||||
@@ -402,6 +406,16 @@ class APIConnection final : public APIServerConnectionBase {
|
||||
void process_state_subscriptions_();
|
||||
#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
|
||||
bool send_message_(uint32_t payload_size, uint16_t message_type, MessageEncodeFn encode_fn, const void *msg);
|
||||
|
||||
@@ -420,7 +434,11 @@ class APIConnection final : public APIServerConnectionBase {
|
||||
// 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.
|
||||
template<typename T> static uint16_t encode_message_to_buffer(T &msg, APIConnection *conn, uint32_t remaining_size) {
|
||||
return encode_to_buffer_slow(T::calc_size_msg(&msg), &T::encode_msg, &msg, conn, remaining_size);
|
||||
if constexpr (T::ESTIMATED_SIZE == 0) {
|
||||
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
|
||||
@@ -432,7 +450,7 @@ class APIConnection final : public APIServerConnectionBase {
|
||||
template<typename T>
|
||||
static uint16_t fill_and_encode_entity_state(EntityBase *entity, T &msg, APIConnection *conn,
|
||||
uint32_t remaining_size) {
|
||||
return fill_and_encode_entity_state(entity, msg, &T::calc_size_msg, &T::encode_msg, conn, remaining_size);
|
||||
return fill_and_encode_entity_state(entity, msg, &calc_size<T>, &proto_encode_msg<T>, conn, remaining_size);
|
||||
}
|
||||
|
||||
// Non-template core — fills info fields, allocates buffers, and encodes
|
||||
@@ -444,7 +462,7 @@ class APIConnection final : public APIServerConnectionBase {
|
||||
template<typename T>
|
||||
static uint16_t fill_and_encode_entity_info(EntityBase *entity, T &msg, APIConnection *conn,
|
||||
uint32_t remaining_size) {
|
||||
return fill_and_encode_entity_info(entity, msg, &T::calc_size_msg, &T::encode_msg, conn, remaining_size);
|
||||
return fill_and_encode_entity_info(entity, msg, &calc_size<T>, &proto_encode_msg<T>, conn, remaining_size);
|
||||
}
|
||||
|
||||
// Non-template core — fills device_class, then delegates to fill_and_encode_entity_info
|
||||
@@ -458,8 +476,8 @@ class APIConnection final : public APIServerConnectionBase {
|
||||
static uint16_t fill_and_encode_entity_info_with_device_class(EntityBase *entity, T &msg,
|
||||
StringRef &device_class_field, APIConnection *conn,
|
||||
uint32_t remaining_size) {
|
||||
return fill_and_encode_entity_info_with_device_class(entity, msg, device_class_field, &T::calc_size_msg,
|
||||
&T::encode_msg, conn, remaining_size);
|
||||
return fill_and_encode_entity_info_with_device_class(entity, msg, device_class_field, &calc_size<T>,
|
||||
&proto_encode_msg<T>, conn, remaining_size);
|
||||
}
|
||||
|
||||
#ifdef USE_VOICE_ASSISTANT
|
||||
|
||||
@@ -46,13 +46,7 @@ inline uint16_t ESPHOME_ALWAYS_INLINE APIConnection::encode_to_buffer(uint32_t c
|
||||
return 0;
|
||||
}
|
||||
ProtoWriteBuffer buffer{&shared_buf, shared_buf.size() - calculated_size};
|
||||
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
|
||||
encode_fn(msg, buffer PROTO_ENCODE_DEBUG_INIT(&shared_buf));
|
||||
|
||||
return total_calculated_size;
|
||||
}
|
||||
|
||||
+2628
-2440
File diff suppressed because it is too large
Load Diff
+305
-654
File diff suppressed because it is too large
Load Diff
@@ -214,74 +214,73 @@ void ProtoDecodableMessage::decode(const uint8_t *buffer, size_t length) {
|
||||
const uint8_t *ptr = buffer;
|
||||
const uint8_t *end = buffer + length;
|
||||
|
||||
// Single-byte varints dominate, so that case advances the cursor inline.
|
||||
auto read_varint = [&](proto_varint_value_t &value) ESPHOME_ALWAYS_INLINE {
|
||||
if (ptr == end)
|
||||
return false;
|
||||
if (*ptr < 0x80) [[likely]] {
|
||||
value = *ptr++;
|
||||
return true;
|
||||
}
|
||||
auto res = ProtoVarInt::parse_non_empty(ptr, end - ptr);
|
||||
if (!res.has_value())
|
||||
return false;
|
||||
value = res.value;
|
||||
ptr += res.consumed;
|
||||
return true;
|
||||
};
|
||||
|
||||
while (ptr < end) {
|
||||
proto_varint_value_t tag_value;
|
||||
if (!read_varint(tag_value)) {
|
||||
// Parse field header - ptr < end guarantees len >= 1
|
||||
auto res = ProtoVarInt::parse_non_empty(ptr, end - ptr);
|
||||
if (!res.has_value()) {
|
||||
ESP_LOGV(TAG, "Invalid field start at offset %ld", (long) (ptr - buffer));
|
||||
return;
|
||||
}
|
||||
|
||||
uint32_t tag = static_cast<uint32_t>(tag_value);
|
||||
uint32_t tag = static_cast<uint32_t>(res.value);
|
||||
uint32_t field_type = tag & WIRE_TYPE_MASK;
|
||||
// Length-delimited fields move this past the length prefix
|
||||
const uint8_t *data = ptr;
|
||||
proto_varint_value_t scalar;
|
||||
uint32_t field_id = tag >> 3;
|
||||
ptr += res.consumed;
|
||||
|
||||
if (field_type == WIRE_TYPE_VARINT) [[likely]] {
|
||||
if (!read_varint(scalar)) {
|
||||
ESP_LOGV(TAG, "Invalid VarInt at offset %ld", (long) (ptr - buffer));
|
||||
return;
|
||||
}
|
||||
} else {
|
||||
switch (field_type) {
|
||||
case WIRE_TYPE_LENGTH_DELIMITED: {
|
||||
proto_varint_value_t length_value;
|
||||
if (!read_varint(length_value)) {
|
||||
ESP_LOGV(TAG, "Invalid Length Delimited at offset %ld", (long) (ptr - buffer));
|
||||
return;
|
||||
}
|
||||
uint32_t field_length = static_cast<uint32_t>(length_value);
|
||||
if (field_length > static_cast<size_t>(end - ptr)) {
|
||||
ESP_LOGV(TAG, "Out-of-bounds Length Delimited at offset %ld", (long) (ptr - buffer));
|
||||
return;
|
||||
}
|
||||
data = ptr;
|
||||
scalar = field_length;
|
||||
ptr += field_length;
|
||||
break;
|
||||
}
|
||||
case WIRE_TYPE_FIXED32: {
|
||||
if (end - ptr < 4) {
|
||||
ESP_LOGV(TAG, "Out-of-bounds Fixed32-bit at offset %ld", (long) (ptr - buffer));
|
||||
return;
|
||||
}
|
||||
// Byte loads instead of memcpy: ESP-IDF passes -fno-builtin-memcpy, which made this a call
|
||||
scalar = encode_uint32(ptr[3], ptr[2], ptr[1], ptr[0]);
|
||||
ptr += 4;
|
||||
break;
|
||||
}
|
||||
default:
|
||||
ESP_LOGV(TAG, "Invalid field type %" PRIu32 " at offset %ld", field_type, (long) (ptr - buffer));
|
||||
switch (field_type) {
|
||||
case WIRE_TYPE_VARINT: { // VarInt
|
||||
res = ProtoVarInt::parse(ptr, end - ptr);
|
||||
if (!res.has_value()) {
|
||||
ESP_LOGV(TAG, "Invalid VarInt at offset %ld", (long) (ptr - buffer));
|
||||
return;
|
||||
}
|
||||
if (!this->decode_varint(field_id, res.value)) {
|
||||
ESP_LOGV(TAG, "Cannot decode VarInt field %" PRIu32 " with value %" PRIu64 "!", field_id,
|
||||
static_cast<uint64_t>(res.value));
|
||||
}
|
||||
ptr += res.consumed;
|
||||
break;
|
||||
}
|
||||
case WIRE_TYPE_LENGTH_DELIMITED: { // Length-delimited
|
||||
res = ProtoVarInt::parse(ptr, end - ptr);
|
||||
if (!res.has_value()) {
|
||||
ESP_LOGV(TAG, "Invalid Length Delimited at offset %ld", (long) (ptr - buffer));
|
||||
return;
|
||||
}
|
||||
uint32_t field_length = static_cast<uint32_t>(res.value);
|
||||
ptr += res.consumed;
|
||||
if (field_length > static_cast<size_t>(end - ptr)) {
|
||||
ESP_LOGV(TAG, "Out-of-bounds Length Delimited at offset %ld", (long) (ptr - buffer));
|
||||
return;
|
||||
}
|
||||
if (!this->decode_length(field_id, ProtoLengthDelimited(ptr, field_length))) {
|
||||
ESP_LOGV(TAG, "Cannot decode Length Delimited field %" PRIu32 "!", field_id);
|
||||
}
|
||||
ptr += field_length;
|
||||
break;
|
||||
}
|
||||
case WIRE_TYPE_FIXED32: { // 32-bit
|
||||
if (end - ptr < 4) {
|
||||
ESP_LOGV(TAG, "Out-of-bounds Fixed32-bit at offset %ld", (long) (ptr - buffer));
|
||||
return;
|
||||
}
|
||||
uint32_t val;
|
||||
#if __BYTE_ORDER__ == __ORDER_LITTLE_ENDIAN__
|
||||
// Protobuf fixed32 is little-endian — direct load on LE platforms
|
||||
memcpy(&val, ptr, 4);
|
||||
#else
|
||||
val = encode_uint32(ptr[3], ptr[2], ptr[1], ptr[0]);
|
||||
#endif
|
||||
if (!this->decode_32bit(field_id, Proto32Bit(val))) {
|
||||
ESP_LOGV(TAG, "Cannot decode 32-bit field %" PRIu32 " with value %" PRIu32 "!", field_id, val);
|
||||
}
|
||||
ptr += 4;
|
||||
break;
|
||||
}
|
||||
default:
|
||||
ESP_LOGV(TAG, "Invalid field type %" PRIu32 " at offset %ld", field_type, (long) (ptr - buffer));
|
||||
return;
|
||||
}
|
||||
this->decode_field(tag, data, scalar);
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+166
-228
@@ -170,43 +170,40 @@ class ProtoVarInt {
|
||||
class ProtoMessage;
|
||||
class ProtoSize;
|
||||
|
||||
/// Case label for decode_field(): the wire tag of a field, so a field that arrives with another wire
|
||||
/// type matches no case.
|
||||
constexpr uint32_t proto_tag(uint32_t field_id, uint32_t wire_type) { return (field_id << 3) | wire_type; }
|
||||
|
||||
/// One decoded field: the payload pointer and a scalar holding the varint or fixed32 value, or the
|
||||
/// length of a length-delimited field. The wire type in the tag says which applies; accessors do not check.
|
||||
class ProtoFieldValue {
|
||||
class ProtoLengthDelimited {
|
||||
public:
|
||||
ProtoFieldValue(const uint8_t *data, proto_varint_value_t scalar) : data_(data), scalar_(scalar) {}
|
||||
explicit ProtoLengthDelimited(const uint8_t *value, size_t length) : value_(value), length_(length) {}
|
||||
std::string as_string() const { return std::string(reinterpret_cast<const char *>(this->value_), this->length_); }
|
||||
|
||||
proto_varint_value_t as_varint() const { return this->scalar_; }
|
||||
// A bool is sent as 0 or 1, so the low word is enough and saves a second compare with 64 bit varints
|
||||
bool as_bool() const { return static_cast<uint32_t>(this->scalar_) != 0; }
|
||||
// Direct access to raw data without string allocation
|
||||
const uint8_t *data() const { return this->value_; }
|
||||
size_t size() const { return this->length_; }
|
||||
|
||||
// Length-delimited accessors
|
||||
const uint8_t *data() const { return this->data_; }
|
||||
size_t size() const { return static_cast<size_t>(this->scalar_); }
|
||||
std::string as_string() const { return std::string(reinterpret_cast<const char *>(this->data_), this->size()); }
|
||||
/// Decode the length-delimited payload into a message instance.
|
||||
/// Decode the length-delimited data into a message instance.
|
||||
/// Template preserves concrete type so decode() resolves statically.
|
||||
template<typename T> void decode_to_message(T &msg) const { msg.decode(this->data_, this->size()); }
|
||||
template<typename T> void decode_to_message(T &msg) const;
|
||||
|
||||
// Fixed32 accessors
|
||||
uint32_t as_fixed32() const { return static_cast<uint32_t>(this->scalar_); }
|
||||
int32_t as_sfixed32() const { return static_cast<int32_t>(this->as_fixed32()); }
|
||||
protected:
|
||||
const uint8_t *const value_;
|
||||
const size_t length_;
|
||||
};
|
||||
|
||||
class Proto32Bit {
|
||||
public:
|
||||
explicit Proto32Bit(uint32_t value) : value_(value) {}
|
||||
uint32_t as_fixed32() const { return this->value_; }
|
||||
int32_t as_sfixed32() const { return static_cast<int32_t>(this->value_); }
|
||||
float as_float() const {
|
||||
union {
|
||||
uint32_t raw;
|
||||
float value;
|
||||
} s{};
|
||||
s.raw = this->as_fixed32();
|
||||
s.raw = this->value_;
|
||||
return s.value;
|
||||
}
|
||||
|
||||
private:
|
||||
const uint8_t *data_;
|
||||
proto_varint_value_t scalar_;
|
||||
protected:
|
||||
const uint32_t value_;
|
||||
};
|
||||
|
||||
// NOTE: Proto64Bit class removed - wire type 1 (64-bit fixed) not supported
|
||||
@@ -255,7 +252,7 @@ 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)); }
|
||||
void encode_field_raw(uint32_t field_id, uint32_t type) { this->encode_varint_raw((field_id << 3) | 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);
|
||||
@@ -290,31 +287,19 @@ class ProtoWriteBuffer {
|
||||
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.
|
||||
constexpr uint32_t VARINT_MAX_1_BYTE = 1 << 7; // 128
|
||||
constexpr uint32_t VARINT_MAX_2_BYTE = 1 << 14; // 16384
|
||||
|
||||
/// 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.
|
||||
/// Static encode helpers for generated encode() functions.
|
||||
/// Generated code hoists buffer.pos_ into a local uint8_t *__restrict__ pos,
|
||||
/// then calls these methods which take pos by reference. No struct, no overhead.
|
||||
/// For sub-messages, pos is synced back to buffer before the call and reloaded after.
|
||||
class ProtoEncode {
|
||||
public:
|
||||
/// Write a multi-byte varint directly through a pos pointer.
|
||||
template<typename T>
|
||||
[[nodiscard]] static inline uint8_t *encode_varint_raw_loop(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
|
||||
T value) {
|
||||
static inline void encode_varint_raw_loop(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, T value) {
|
||||
do {
|
||||
PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
|
||||
*pos++ = static_cast<uint8_t>(value | 0x80);
|
||||
@@ -322,49 +307,48 @@ class ProtoEncode {
|
||||
} while (value > 0x7F);
|
||||
PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
|
||||
*pos++ = static_cast<uint8_t>(value);
|
||||
return pos;
|
||||
}
|
||||
[[nodiscard]] static inline uint8_t *ESPHOME_ALWAYS_INLINE
|
||||
encode_varint_raw(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint32_t value) {
|
||||
static inline void ESPHOME_ALWAYS_INLINE encode_varint_raw(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM,
|
||||
uint32_t value) {
|
||||
if (value < VARINT_MAX_1_BYTE) [[likely]] {
|
||||
PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
|
||||
*pos++ = static_cast<uint8_t>(value);
|
||||
return pos;
|
||||
return;
|
||||
}
|
||||
return encode_varint_raw_loop(pos PROTO_ENCODE_DEBUG_ARG, value);
|
||||
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).
|
||||
[[nodiscard]] static inline uint8_t *ESPHOME_ALWAYS_INLINE
|
||||
encode_varint_raw_short(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint32_t value) {
|
||||
static inline void ESPHOME_ALWAYS_INLINE encode_varint_raw_short(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM,
|
||||
uint32_t value) {
|
||||
if (value < VARINT_MAX_1_BYTE) [[likely]] {
|
||||
PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
|
||||
*pos++ = static_cast<uint8_t>(value);
|
||||
return pos;
|
||||
return;
|
||||
}
|
||||
if (value < VARINT_MAX_2_BYTE) [[likely]] {
|
||||
PROTO_ENCODE_CHECK_BOUNDS(pos, 2);
|
||||
*pos++ = static_cast<uint8_t>(value | 0x80);
|
||||
*pos++ = static_cast<uint8_t>(value >> 7);
|
||||
return pos;
|
||||
return;
|
||||
}
|
||||
return encode_varint_raw_loop(pos PROTO_ENCODE_DEBUG_ARG, value);
|
||||
encode_varint_raw_loop(pos PROTO_ENCODE_DEBUG_ARG, value);
|
||||
}
|
||||
[[nodiscard]] static inline uint8_t *ESPHOME_ALWAYS_INLINE
|
||||
encode_varint_raw_64(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint64_t value) {
|
||||
static inline void ESPHOME_ALWAYS_INLINE encode_varint_raw_64(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM,
|
||||
uint64_t value) {
|
||||
if (value < VARINT_MAX_1_BYTE) [[likely]] {
|
||||
PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
|
||||
*pos++ = static_cast<uint8_t>(value);
|
||||
return pos;
|
||||
return;
|
||||
}
|
||||
return encode_varint_raw_loop(pos PROTO_ENCODE_DEBUG_ARG, value);
|
||||
encode_varint_raw_loop(pos PROTO_ENCODE_DEBUG_ARG, value);
|
||||
}
|
||||
/// Encode a 48-bit MAC address (stored in a uint64) as varint.
|
||||
/// Real MAC addresses occupy the full 48 bits (OUI in upper 24), so the
|
||||
/// fast path -- any non-zero bit in the top 6 of 48 -- emits exactly 7 bytes
|
||||
/// with no per-byte branch. Falls back to the general loop otherwise.
|
||||
/// Caller must guarantee value fits in 48 bits (checked in debug builds).
|
||||
[[nodiscard]] static inline uint8_t *ESPHOME_ALWAYS_INLINE
|
||||
encode_varint_raw_48bit(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint64_t value) {
|
||||
static inline void ESPHOME_ALWAYS_INLINE encode_varint_raw_48bit(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM,
|
||||
uint64_t value) {
|
||||
#ifdef ESPHOME_DEBUG_API
|
||||
assert(value < (1ULL << (MAC_ADDRESS_SIZE * 8)) && "encode_varint_raw_48bit: value exceeds 48 bits");
|
||||
#endif
|
||||
@@ -379,39 +363,38 @@ class ProtoEncode {
|
||||
pos[4] = static_cast<uint8_t>((value >> 28) | 0x80);
|
||||
pos[5] = static_cast<uint8_t>((value >> 35) | 0x80);
|
||||
pos[6] = static_cast<uint8_t>(value >> 42);
|
||||
return pos + 7;
|
||||
pos += 7;
|
||||
return;
|
||||
}
|
||||
return encode_varint_raw_64(pos PROTO_ENCODE_DEBUG_ARG, value);
|
||||
encode_varint_raw_64(pos PROTO_ENCODE_DEBUG_ARG, value);
|
||||
}
|
||||
[[nodiscard]] static inline uint8_t *ESPHOME_ALWAYS_INLINE
|
||||
encode_field_raw(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id, uint32_t type) {
|
||||
return encode_varint_raw(pos PROTO_ENCODE_DEBUG_ARG, proto_tag(field_id, type));
|
||||
static inline void ESPHOME_ALWAYS_INLINE encode_field_raw(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM,
|
||||
uint32_t field_id, uint32_t type) {
|
||||
encode_varint_raw(pos PROTO_ENCODE_DEBUG_ARG, (field_id << 3) | type);
|
||||
}
|
||||
/// Write a single precomputed tag byte. Tag must be < 128.
|
||||
[[nodiscard]] static inline uint8_t *ESPHOME_ALWAYS_INLINE
|
||||
write_raw_byte(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint8_t b) {
|
||||
static inline void ESPHOME_ALWAYS_INLINE write_raw_byte(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM,
|
||||
uint8_t b) {
|
||||
PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
|
||||
*pos++ = b;
|
||||
return pos;
|
||||
}
|
||||
/// Reserve one byte for later backpatch (e.g., sub-message length).
|
||||
/// Advances pos past the reserved byte without writing a value.
|
||||
[[nodiscard]] static inline uint8_t *ESPHOME_ALWAYS_INLINE
|
||||
reserve_byte(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM) {
|
||||
static inline void ESPHOME_ALWAYS_INLINE reserve_byte(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM) {
|
||||
PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
|
||||
return pos + 1;
|
||||
pos++;
|
||||
}
|
||||
/// Write raw bytes to the buffer (no tag, no length prefix).
|
||||
[[nodiscard]] static inline uint8_t *ESPHOME_ALWAYS_INLINE
|
||||
encode_raw(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, const void *data, size_t len) {
|
||||
static inline void ESPHOME_ALWAYS_INLINE encode_raw(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM,
|
||||
const void *data, size_t len) {
|
||||
PROTO_ENCODE_CHECK_BOUNDS(pos, len);
|
||||
std::memcpy(pos, data, len);
|
||||
return pos + len;
|
||||
pos += len;
|
||||
}
|
||||
/// Encode tag + 1-byte length + raw string data. For strings with max_data_length < 128.
|
||||
/// Tag must be a single-byte varint (< 128). Always encodes (no zero check).
|
||||
[[nodiscard]] static inline uint8_t *encode_short_string_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
|
||||
uint8_t tag, const StringRef &ref) {
|
||||
static inline void encode_short_string_force(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint8_t tag,
|
||||
const StringRef &ref) {
|
||||
#ifdef ESPHOME_DEBUG_API
|
||||
assert(ref.size() < 128 && "encode_short_string_force: string exceeds max_data_length < 128");
|
||||
#endif
|
||||
@@ -419,191 +402,137 @@ class ProtoEncode {
|
||||
pos[0] = tag;
|
||||
pos[1] = static_cast<uint8_t>(ref.size());
|
||||
std::memcpy(pos + 2, ref.c_str(), ref.size());
|
||||
return pos + 2 + ref.size();
|
||||
pos += 2 + ref.size();
|
||||
}
|
||||
/// Write a precomputed tag byte + 32-bit value. Outlined on embedded: one copy beats inline stores per field.
|
||||
[[nodiscard]] static PROTO_OUTLINE_FOR_SIZE uint8_t *write_tag_and_fixed32(
|
||||
uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM, uint8_t tag, uint32_t value) {
|
||||
/// Write a precomputed tag byte + 32-bit value in one operation.
|
||||
static inline void ESPHOME_ALWAYS_INLINE write_tag_and_fixed32(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM,
|
||||
uint8_t tag, uint32_t value) {
|
||||
PROTO_ENCODE_CHECK_BOUNDS(pos, 5);
|
||||
pos[0] = tag;
|
||||
write_fixed32_le(pos + 1, value);
|
||||
return pos + 5;
|
||||
#if __BYTE_ORDER__ == __ORDER_LITTLE_ENDIAN__
|
||||
std::memcpy(pos + 1, &value, 4);
|
||||
#else
|
||||
pos[1] = static_cast<uint8_t>(value & 0xFF);
|
||||
pos[2] = static_cast<uint8_t>((value >> 8) & 0xFF);
|
||||
pos[3] = static_cast<uint8_t>((value >> 16) & 0xFF);
|
||||
pos[4] = static_cast<uint8_t>((value >> 24) & 0xFF);
|
||||
#endif
|
||||
pos += 5;
|
||||
}
|
||||
[[nodiscard]] static inline uint8_t *encode_string_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
|
||||
uint32_t field_id, const char *string, size_t len) {
|
||||
pos = encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 2); // type 2: Length-delimited string
|
||||
static inline void encode_string(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
|
||||
const char *string, size_t len, bool force = false) {
|
||||
if (len == 0 && !force)
|
||||
return;
|
||||
encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 2); // type 2: Length-delimited string
|
||||
// NOLINTNEXTLINE(readability-inconsistent-ifelse-braces) -- false positive on [[likely]] attribute
|
||||
if (len < VARINT_MAX_1_BYTE) [[likely]] {
|
||||
PROTO_ENCODE_CHECK_BOUNDS(pos, 1 + len);
|
||||
*pos++ = static_cast<uint8_t>(len);
|
||||
} else {
|
||||
pos = encode_varint_raw_loop(pos PROTO_ENCODE_DEBUG_ARG, len);
|
||||
encode_varint_raw_loop(pos PROTO_ENCODE_DEBUG_ARG, len);
|
||||
PROTO_ENCODE_CHECK_BOUNDS(pos, len);
|
||||
}
|
||||
std::memcpy(pos, string, len);
|
||||
return pos + len;
|
||||
pos += len;
|
||||
}
|
||||
[[nodiscard]] static inline uint8_t *encode_string(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
|
||||
uint32_t field_id, const char *string, size_t len) {
|
||||
if (len == 0)
|
||||
return pos;
|
||||
return encode_string_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, string, len);
|
||||
static inline void encode_string(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
|
||||
const std::string &value, bool force = false) {
|
||||
encode_string(pos PROTO_ENCODE_DEBUG_ARG, field_id, value.data(), value.size(), force);
|
||||
}
|
||||
[[nodiscard]] static inline uint8_t *encode_string_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
|
||||
uint32_t field_id, const std::string &value) {
|
||||
return encode_string_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, value.data(), value.size());
|
||||
static inline void encode_string(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
|
||||
const StringRef &ref, bool force = false) {
|
||||
encode_string(pos PROTO_ENCODE_DEBUG_ARG, field_id, ref.c_str(), ref.size(), force);
|
||||
}
|
||||
[[nodiscard]] static inline uint8_t *encode_string(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
|
||||
uint32_t field_id, const StringRef &ref) {
|
||||
return encode_string(pos PROTO_ENCODE_DEBUG_ARG, field_id, ref.c_str(), ref.size());
|
||||
static inline void encode_bytes(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
|
||||
const uint8_t *data, size_t len, bool force = false) {
|
||||
encode_string(pos PROTO_ENCODE_DEBUG_ARG, field_id, reinterpret_cast<const char *>(data), len, force);
|
||||
}
|
||||
[[nodiscard]] static inline uint8_t *encode_string_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
|
||||
uint32_t field_id, const StringRef &ref) {
|
||||
return encode_string_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, ref.c_str(), ref.size());
|
||||
static inline void encode_uint32(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
|
||||
uint32_t value, bool force = false) {
|
||||
if (value == 0 && !force)
|
||||
return;
|
||||
encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 0);
|
||||
encode_varint_raw(pos PROTO_ENCODE_DEBUG_ARG, value);
|
||||
}
|
||||
[[nodiscard]] static inline uint8_t *encode_bytes(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
|
||||
uint32_t field_id, const uint8_t *data, size_t len) {
|
||||
return encode_string(pos PROTO_ENCODE_DEBUG_ARG, field_id, reinterpret_cast<const char *>(data), len);
|
||||
static inline void encode_uint64(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
|
||||
uint64_t value, bool force = false) {
|
||||
if (value == 0 && !force)
|
||||
return;
|
||||
encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 0);
|
||||
encode_varint_raw_64(pos PROTO_ENCODE_DEBUG_ARG, value);
|
||||
}
|
||||
[[nodiscard]] static inline uint8_t *encode_bytes_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
|
||||
uint32_t field_id, const uint8_t *data, size_t len) {
|
||||
return encode_string_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, reinterpret_cast<const char *>(data), len);
|
||||
}
|
||||
[[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);
|
||||
static inline void encode_bool(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id, bool value,
|
||||
bool force = false) {
|
||||
if (!value && !force)
|
||||
return;
|
||||
encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 0);
|
||||
PROTO_ENCODE_CHECK_BOUNDS(pos, 1);
|
||||
*pos++ = value ? 0x01 : 0x00;
|
||||
return pos;
|
||||
}
|
||||
[[nodiscard]] static inline uint8_t *encode_bool(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
|
||||
uint32_t field_id, bool value) {
|
||||
if (!value)
|
||||
return pos;
|
||||
return encode_bool_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, value);
|
||||
}
|
||||
/// 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);
|
||||
static inline void encode_fixed32(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
|
||||
uint32_t value, bool force = false) {
|
||||
if (value == 0 && !force)
|
||||
return;
|
||||
encode_field_raw(pos PROTO_ENCODE_DEBUG_ARG, field_id, 5);
|
||||
PROTO_ENCODE_CHECK_BOUNDS(pos, 4);
|
||||
write_fixed32_le(pos, value);
|
||||
return pos + 4;
|
||||
}
|
||||
[[nodiscard]] static inline uint8_t *encode_fixed32(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
|
||||
uint32_t field_id, uint32_t value) {
|
||||
if (value == 0)
|
||||
return pos;
|
||||
return encode_fixed32_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, value);
|
||||
#if __BYTE_ORDER__ == __ORDER_LITTLE_ENDIAN__
|
||||
std::memcpy(pos, &value, 4);
|
||||
pos += 4;
|
||||
#else
|
||||
*pos++ = (value >> 0) & 0xFF;
|
||||
*pos++ = (value >> 8) & 0xFF;
|
||||
*pos++ = (value >> 16) & 0xFF;
|
||||
*pos++ = (value >> 24) & 0xFF;
|
||||
#endif
|
||||
}
|
||||
// NOTE: Wire type 1 (64-bit fixed: double, fixed64, sfixed64) is intentionally
|
||||
// not supported to reduce overhead on embedded systems. All ESPHome devices are
|
||||
// 32-bit microcontrollers where 64-bit operations are expensive. If 64-bit support
|
||||
// is needed in the future, the necessary encoding/decoding functions must be added.
|
||||
[[nodiscard]] static inline uint8_t *encode_float(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
|
||||
uint32_t field_id, float value) {
|
||||
return encode_fixed32(pos PROTO_ENCODE_DEBUG_ARG, field_id, float_to_raw(value));
|
||||
static inline void encode_float(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id, float value,
|
||||
bool force = false) {
|
||||
uint32_t raw = float_to_raw(value);
|
||||
if (raw == 0 && !force)
|
||||
return;
|
||||
encode_fixed32(pos PROTO_ENCODE_DEBUG_ARG, field_id, raw);
|
||||
}
|
||||
[[nodiscard]] static inline uint8_t *encode_float_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
|
||||
uint32_t field_id, float value) {
|
||||
return encode_fixed32_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, float_to_raw(value));
|
||||
}
|
||||
[[nodiscard]] static inline uint8_t *encode_int32_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
|
||||
uint32_t field_id, int32_t value) {
|
||||
static inline void encode_int32(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id, int32_t value,
|
||||
bool force = false) {
|
||||
if (value < 0) {
|
||||
// negative int32 is always 10 byte long
|
||||
return encode_uint64_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, static_cast<uint64_t>(value));
|
||||
encode_uint64(pos PROTO_ENCODE_DEBUG_ARG, field_id, static_cast<uint64_t>(value), force);
|
||||
return;
|
||||
}
|
||||
return encode_uint32_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, static_cast<uint32_t>(value));
|
||||
encode_uint32(pos PROTO_ENCODE_DEBUG_ARG, field_id, static_cast<uint32_t>(value), force);
|
||||
}
|
||||
[[nodiscard]] static inline uint8_t *encode_int32(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
|
||||
uint32_t field_id, int32_t value) {
|
||||
if (value == 0)
|
||||
return pos;
|
||||
return encode_int32_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, value);
|
||||
static inline void encode_int64(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id, int64_t value,
|
||||
bool force = false) {
|
||||
encode_uint64(pos PROTO_ENCODE_DEBUG_ARG, field_id, static_cast<uint64_t>(value), force);
|
||||
}
|
||||
[[nodiscard]] static inline uint8_t *encode_int64(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
|
||||
uint32_t field_id, int64_t value) {
|
||||
return encode_uint64(pos PROTO_ENCODE_DEBUG_ARG, field_id, static_cast<uint64_t>(value));
|
||||
static inline void encode_sint32(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
|
||||
int32_t value, bool force = false) {
|
||||
encode_uint32(pos PROTO_ENCODE_DEBUG_ARG, field_id, encode_zigzag32(value), force);
|
||||
}
|
||||
[[nodiscard]] static inline uint8_t *encode_int64_force(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
|
||||
uint32_t field_id, int64_t value) {
|
||||
return encode_uint64_force(pos PROTO_ENCODE_DEBUG_ARG, field_id, static_cast<uint64_t>(value));
|
||||
static inline void encode_sint64(uint8_t *__restrict__ &pos PROTO_ENCODE_DEBUG_PARAM, uint32_t field_id,
|
||||
int64_t value, bool force = false) {
|
||||
encode_uint64(pos PROTO_ENCODE_DEBUG_ARG, field_id, encode_zigzag64(value), force);
|
||||
}
|
||||
[[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.
|
||||
/// Sub-message encoding: sync pos to buffer, delegate, get pos from return value.
|
||||
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) {
|
||||
static inline void 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();
|
||||
pos = buffer.get_pos();
|
||||
}
|
||||
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) {
|
||||
static inline void 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();
|
||||
}
|
||||
|
||||
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);
|
||||
}
|
||||
pos = buffer.get_pos();
|
||||
}
|
||||
};
|
||||
#undef PROTO_OUTLINE_FOR_SIZE
|
||||
#undef PROTO_FIXED32_BYTE_STORES
|
||||
|
||||
#ifdef HAS_PROTO_MESSAGE_DUMP
|
||||
/**
|
||||
@@ -695,12 +624,11 @@ class DumpBuffer {
|
||||
|
||||
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 uint32_t calc_size_msg(const void *self) { return 0; }
|
||||
// Non-virtual defaults for messages with no fields.
|
||||
// Concrete message classes hide these with their own implementations.
|
||||
// All call sites use templates to preserve the concrete type, so virtual
|
||||
// dispatch is not needed. This eliminates per-message vtable entries for
|
||||
// encode/calculate_size, saving ~1.3 KB of flash across all message types.
|
||||
uint8_t *encode(ProtoWriteBuffer &buffer PROTO_ENCODE_DEBUG_PARAM) const { return buffer.get_pos(); }
|
||||
uint32_t calculate_size() const { return 0; }
|
||||
#ifdef HAS_PROTO_MESSAGE_DUMP
|
||||
@@ -735,10 +663,10 @@ class ProtoDecodableMessage : public ProtoMessage {
|
||||
|
||||
protected:
|
||||
~ProtoDecodableMessage() = default;
|
||||
/// Store one decoded field; \p scalar is the varint or fixed32 value, or the length of the
|
||||
/// length-delimited payload at \p data. An unknown field or wrong wire type matches no case and is skipped.
|
||||
/// Three register arguments keep the decode loop free of spills.
|
||||
virtual void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) {}
|
||||
virtual bool decode_varint(uint32_t field_id, proto_varint_value_t value) { return false; }
|
||||
virtual bool decode_length(uint32_t field_id, ProtoLengthDelimited value) { return false; }
|
||||
virtual bool decode_32bit(uint32_t field_id, Proto32Bit value) { return false; }
|
||||
// NOTE: decode_64bit removed - wire type 1 not supported
|
||||
};
|
||||
|
||||
class ProtoSize {
|
||||
@@ -864,7 +792,7 @@ class ProtoSize {
|
||||
* @return The number of bytes needed to encode the field ID and wire type
|
||||
*/
|
||||
static constexpr uint32_t field(uint32_t field_id, uint32_t type) {
|
||||
uint32_t tag = proto_tag(field_id, type & WIRE_TYPE_MASK);
|
||||
uint32_t tag = (field_id << 3) | (type & WIRE_TYPE_MASK);
|
||||
return varint(tag);
|
||||
}
|
||||
|
||||
@@ -948,14 +876,24 @@ class ProtoSize {
|
||||
|
||||
// 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.
|
||||
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);
|
||||
this->encode_sub_message(field_id, &value, &proto_encode_msg<T>);
|
||||
}
|
||||
|
||||
// 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);
|
||||
this->encode_optional_sub_message(field_id, value.calculate_size(), &value, &proto_encode_msg<T>);
|
||||
}
|
||||
|
||||
// Template decode_to_message - preserves concrete type so decode() resolves statically
|
||||
template<typename T> void ProtoLengthDelimited::decode_to_message(T &msg) const {
|
||||
msg.decode(this->value_, this->length_);
|
||||
}
|
||||
|
||||
template<typename T> const char *proto_enum_to_string(T value);
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import esphome.codegen as cg
|
||||
from esphome.components import climate_ir
|
||||
from esphome.components import climate_ir, remote_base
|
||||
from esphome.types import ConfigType
|
||||
|
||||
AUTO_LOAD = ["climate_ir"]
|
||||
@@ -12,4 +12,5 @@ CONFIG_SCHEMA = climate_ir.climate_ir_with_receiver_schema(CoolixClimate)
|
||||
|
||||
|
||||
async def to_code(config: ConfigType) -> None:
|
||||
remote_base.request_protocol("coolix") # used from C++
|
||||
await climate_ir.new_climate_ir(config)
|
||||
|
||||
@@ -59,11 +59,6 @@ void Infrared::setup() {
|
||||
// Set up traits based on configuration
|
||||
this->traits_.set_supports_transmitter(this->has_transmitter());
|
||||
this->traits_.set_supports_receiver(this->has_receiver());
|
||||
|
||||
// Register as listener for received IR data
|
||||
if (this->receiver_ != nullptr) {
|
||||
this->receiver_->register_listener(this);
|
||||
}
|
||||
}
|
||||
|
||||
void Infrared::dump_config() {
|
||||
|
||||
@@ -1,6 +1,9 @@
|
||||
"""IR/RF Proxy component - provides remote_base backend for infrared platform."""
|
||||
|
||||
import esphome.codegen as cg
|
||||
from esphome.components import remote_base
|
||||
from esphome.cpp_generator import MockObj
|
||||
from esphome.types import ConfigType
|
||||
|
||||
CODEOWNERS = ["@kbx81"]
|
||||
|
||||
@@ -9,3 +12,10 @@ ir_rf_proxy_ns = cg.esphome_ns.namespace("ir_rf_proxy")
|
||||
|
||||
CONF_REMOTE_RECEIVER_ID = "remote_receiver_id"
|
||||
CONF_REMOTE_TRANSMITTER_ID = "remote_transmitter_id"
|
||||
|
||||
|
||||
async def attach_receiver(var: MockObj, config: ConfigType) -> None:
|
||||
"""Wire the configured remote_receiver to a proxy entity and register it as a listener."""
|
||||
receiver = await cg.get_variable(config[CONF_REMOTE_RECEIVER_ID])
|
||||
cg.add(var.set_receiver(receiver))
|
||||
remote_base.add_listener(receiver, var)
|
||||
|
||||
@@ -9,7 +9,12 @@ import esphome.config_validation as cv
|
||||
from esphome.const import CONF_CARRIER_DUTY_PERCENT, CONF_FREQUENCY
|
||||
import esphome.final_validate as fv
|
||||
|
||||
from . import CONF_REMOTE_RECEIVER_ID, CONF_REMOTE_TRANSMITTER_ID, ir_rf_proxy_ns
|
||||
from . import (
|
||||
CONF_REMOTE_RECEIVER_ID,
|
||||
CONF_REMOTE_TRANSMITTER_ID,
|
||||
attach_receiver,
|
||||
ir_rf_proxy_ns,
|
||||
)
|
||||
|
||||
CODEOWNERS = ["@kbx81"]
|
||||
DEPENDENCIES = ["infrared"]
|
||||
@@ -82,8 +87,7 @@ async def to_code(config: dict[str, Any]) -> None:
|
||||
|
||||
# Link receiver if specified
|
||||
if CONF_REMOTE_RECEIVER_ID in config:
|
||||
receiver = await cg.get_variable(config[CONF_REMOTE_RECEIVER_ID])
|
||||
cg.add(var.set_receiver(receiver))
|
||||
await attach_receiver(var, config)
|
||||
|
||||
# Set receiver demodulation frequency if specified (metadata only, no hardware effect)
|
||||
if CONF_RECEIVER_FREQUENCY in config:
|
||||
|
||||
@@ -97,10 +97,6 @@ void RfProxy::setup() {
|
||||
|
||||
// remote_transmitter/receiver always uses OOK (on-off keying)
|
||||
this->traits_.add_supported_modulation(radio_frequency::RadioFrequencyModulation::RADIO_FREQUENCY_MODULATION_OOK);
|
||||
|
||||
if (this->receiver_ != nullptr) {
|
||||
this->receiver_->register_listener(this);
|
||||
}
|
||||
}
|
||||
|
||||
void RfProxy::dump_config() {
|
||||
|
||||
@@ -7,7 +7,12 @@ from esphome.const import CONF_CARRIER_DUTY_PERCENT, CONF_FREQUENCY
|
||||
import esphome.final_validate as fv
|
||||
from esphome.types import ConfigType
|
||||
|
||||
from . import CONF_REMOTE_RECEIVER_ID, CONF_REMOTE_TRANSMITTER_ID, ir_rf_proxy_ns
|
||||
from . import (
|
||||
CONF_REMOTE_RECEIVER_ID,
|
||||
CONF_REMOTE_TRANSMITTER_ID,
|
||||
attach_receiver,
|
||||
ir_rf_proxy_ns,
|
||||
)
|
||||
|
||||
CODEOWNERS = ["@kbx81"]
|
||||
DEPENDENCIES = ["radio_frequency"]
|
||||
@@ -66,5 +71,4 @@ async def to_code(config: ConfigType) -> None:
|
||||
cg.add(var.set_transmitter(transmitter))
|
||||
|
||||
if CONF_REMOTE_RECEIVER_ID in config:
|
||||
receiver = await cg.get_variable(config[CONF_REMOTE_RECEIVER_ID])
|
||||
cg.add(var.set_receiver(receiver))
|
||||
await attach_receiver(var, config)
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from esphome import automation
|
||||
import esphome.codegen as cg
|
||||
from esphome.components import climate, remote_transmitter, sensor, uart
|
||||
from esphome.components import climate, remote_base, remote_transmitter, sensor, uart
|
||||
from esphome.components.climate import ClimateMode, ClimatePreset, ClimateSwingMode
|
||||
from esphome.components.remote_base import CONF_TRANSMITTER_ID
|
||||
import esphome.config_validation as cv
|
||||
@@ -280,6 +280,7 @@ async def to_code(config):
|
||||
cg.add(var.set_response_timeout(config[CONF_TIMEOUT].total_milliseconds))
|
||||
cg.add(var.set_request_attempts(config[CONF_NUM_ATTEMPTS]))
|
||||
if CONF_TRANSMITTER_ID in config:
|
||||
remote_base.request_protocol("midea") # ir_transmitter.h uses it from C++
|
||||
cg.add_define("USE_REMOTE_TRANSMITTER")
|
||||
transmitter_ = await cg.get_variable(config[CONF_TRANSMITTER_ID])
|
||||
cg.add(var.set_transmitter(transmitter_))
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import esphome.codegen as cg
|
||||
from esphome.components import climate_ir
|
||||
from esphome.components import climate_ir, remote_base
|
||||
import esphome.config_validation as cv
|
||||
from esphome.const import CONF_USE_FAHRENHEIT
|
||||
from esphome.types import ConfigType
|
||||
@@ -19,5 +19,9 @@ CONFIG_SCHEMA = climate_ir.climate_ir_with_receiver_schema(MideaIR).extend(
|
||||
|
||||
|
||||
async def to_code(config: ConfigType) -> None:
|
||||
# midea_ir uses MideaProtocol from C++ and auto-loads coolix, whose coolix.cpp uses
|
||||
# CoolixProtocol even when no coolix climate is configured
|
||||
remote_base.request_protocol("midea")
|
||||
remote_base.request_protocol("coolix")
|
||||
var = await climate_ir.new_climate_ir(config)
|
||||
cg.add(var.set_fahrenheit(config[CONF_USE_FAHRENHEIT]))
|
||||
|
||||
@@ -1,6 +1,11 @@
|
||||
from collections.abc import Callable
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from esphome import automation
|
||||
import esphome.codegen as cg
|
||||
from esphome.components import binary_sensor
|
||||
from esphome.config_helpers import filter_source_files_from_defines
|
||||
import esphome.config_validation as cv
|
||||
from esphome.const import (
|
||||
CONF_ADDRESS,
|
||||
@@ -40,7 +45,9 @@ from esphome.const import (
|
||||
CONF_ZERO,
|
||||
)
|
||||
from esphome.core import ID, coroutine
|
||||
from esphome.cpp_generator import MockObj
|
||||
from esphome.schema_extractors import SCHEMA_EXTRACT, schema_extractor
|
||||
from esphome.types import ConfigType
|
||||
from esphome.util import Registry, SimpleRegistry
|
||||
|
||||
AUTO_LOAD = ["binary_sensor"]
|
||||
@@ -90,9 +97,25 @@ REMOTE_TRANSMITTABLE_SCHEMA = cv.Schema(
|
||||
)
|
||||
|
||||
|
||||
async def register_listener(var, config):
|
||||
# Listener and dumper lists are StaticVectors sized from these counts, so every
|
||||
# registration must go through add_listener / add_dumper
|
||||
_request_listener_slot = cg.slot_counter("REMOTE_BASE_LISTENER_COUNT")
|
||||
_request_dumper_slot = cg.slot_counter("REMOTE_BASE_DUMPER_COUNT")
|
||||
|
||||
|
||||
def add_listener(receiver: MockObj, listener: MockObj) -> None:
|
||||
_request_listener_slot()
|
||||
cg.add(receiver.register_listener(listener))
|
||||
|
||||
|
||||
def add_dumper(receiver: MockObj, dumper: MockObj) -> None:
|
||||
_request_dumper_slot()
|
||||
cg.add(receiver.register_dumper(dumper))
|
||||
|
||||
|
||||
async def register_listener(var: MockObj, config: ConfigType) -> None:
|
||||
receiver = await cg.get_variable(config[CONF_RECEIVER_ID])
|
||||
cg.add(receiver.register_listener(var))
|
||||
add_listener(receiver, var)
|
||||
|
||||
|
||||
async def register_transmittable(var, config):
|
||||
@@ -100,8 +123,47 @@ async def register_transmittable(var, config):
|
||||
cg.add(var.set_transmitter(transmitter_))
|
||||
|
||||
|
||||
def register_binary_sensor(name, type, schema):
|
||||
return BINARY_SENSOR_REGISTRY.register(name, type, schema)
|
||||
# Registry names that share a protocol source file
|
||||
def _protocol_stem(name: str) -> str:
|
||||
if name.startswith("rc_switch"):
|
||||
return "rc_switch"
|
||||
if name == "canalsatld":
|
||||
return "canalsat"
|
||||
return name
|
||||
|
||||
|
||||
def protocol_define(name: str) -> str:
|
||||
return f"USE_REMOTE_PROTOCOL_{_protocol_stem(name).upper()}"
|
||||
|
||||
|
||||
def request_protocol(name: str) -> None:
|
||||
"""Keep a protocol's source file in the build; components using it from C++ must call this."""
|
||||
cg.add_define(protocol_define(name))
|
||||
|
||||
|
||||
_PROTOCOL_STEMS = sorted(
|
||||
path.name.removesuffix("_protocol.cpp")
|
||||
for path in Path(__file__).parent.glob("*_protocol.cpp")
|
||||
)
|
||||
# Only the protocol sources a configuration uses are compiled
|
||||
FILTER_SOURCE_FILES = filter_source_files_from_defines(
|
||||
{f"{stem}_protocol.cpp": protocol_define(stem) for stem in _PROTOCOL_STEMS}
|
||||
)
|
||||
|
||||
|
||||
def register_binary_sensor(
|
||||
name: str, type: MockObj, schema: cv.Schema | dict
|
||||
) -> Callable[[Callable[[MockObj, ConfigType], Any]], Callable]:
|
||||
registerer = BINARY_SENSOR_REGISTRY.register(name, type, schema)
|
||||
|
||||
def decorator(func: Callable[[MockObj, ConfigType], Any]) -> Callable:
|
||||
async def new_func(var: MockObj, config: ConfigType) -> None:
|
||||
request_protocol(name)
|
||||
await coroutine(func)(var, config)
|
||||
|
||||
return registerer(new_func)
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
def register_trigger(name, type, data_type):
|
||||
@@ -114,6 +176,7 @@ def register_trigger(name, type, data_type):
|
||||
|
||||
def decorator(func):
|
||||
async def new_func(config):
|
||||
request_protocol(name)
|
||||
var = cg.new_Pvariable(config[CONF_TRIGGER_ID])
|
||||
await coroutine(func)(var, config)
|
||||
await automation.build_automation(var, [(data_type, "x")], config)
|
||||
@@ -131,6 +194,7 @@ def register_dumper(name, type, schema=None):
|
||||
|
||||
def decorator(func):
|
||||
async def new_func(config, dumper_id):
|
||||
request_protocol(name)
|
||||
var = cg.new_Pvariable(dumper_id)
|
||||
await coroutine(func)(var, config)
|
||||
return var
|
||||
@@ -171,6 +235,7 @@ def register_action(name, type_, schema):
|
||||
|
||||
def decorator(func):
|
||||
async def new_func(config, action_id, template_arg, args):
|
||||
request_protocol(name)
|
||||
var = cg.new_Pvariable(action_id, template_arg)
|
||||
await register_transmittable(var, config)
|
||||
if CONF_REPEAT in config:
|
||||
@@ -210,9 +275,21 @@ TRIGGER_REGISTRY = SimpleRegistry()
|
||||
DUMPER_REGISTRY = Registry()
|
||||
|
||||
|
||||
def _dumper_key(item: Any) -> Any:
|
||||
"""Registry key of a dump entry in either its string or its mapping form."""
|
||||
if isinstance(item, dict) and len(item) == 1:
|
||||
return next(iter(item))
|
||||
return item
|
||||
|
||||
|
||||
def validate_dumpers(value):
|
||||
if isinstance(value, str) and value.lower() == "all":
|
||||
return validate_dumpers(list(DUMPER_REGISTRY.keys()))
|
||||
if isinstance(value, list):
|
||||
# a dumper listed twice would register twice; the receiver holds one secondary dumper
|
||||
keys = [_dumper_key(item) for item in value]
|
||||
if all(isinstance(key, str) for key in keys):
|
||||
value = list(dict(zip(keys, value, strict=True)).values())
|
||||
return cv.validate_registry("dumper", DUMPER_REGISTRY)(value)
|
||||
|
||||
|
||||
|
||||
@@ -191,9 +191,9 @@ class ABBWelcomeData {
|
||||
|
||||
class ABBWelcomeProtocol : public RemoteProtocol<ABBWelcomeData> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const ABBWelcomeData &src) override;
|
||||
optional<ABBWelcomeData> decode(RemoteReceiveData src) override;
|
||||
void dump(const ABBWelcomeData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const ABBWelcomeData &src);
|
||||
optional<ABBWelcomeData> decode(RemoteReceiveData src);
|
||||
void dump(const ABBWelcomeData &data);
|
||||
|
||||
protected:
|
||||
void encode_byte_(RemoteTransmitData *dst, uint8_t data) const;
|
||||
|
||||
@@ -15,9 +15,9 @@ struct AEHAData {
|
||||
|
||||
class AEHAProtocol : public RemoteProtocol<AEHAData> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const AEHAData &data) override;
|
||||
optional<AEHAData> decode(RemoteReceiveData src) override;
|
||||
void dump(const AEHAData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const AEHAData &data);
|
||||
optional<AEHAData> decode(RemoteReceiveData src);
|
||||
void dump(const AEHAData &data);
|
||||
|
||||
private:
|
||||
std::string format_data_(const std::vector<uint8_t> &data);
|
||||
|
||||
@@ -16,9 +16,9 @@ struct Beo4Data {
|
||||
|
||||
class Beo4Protocol : public RemoteProtocol<Beo4Data> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const Beo4Data &data) override;
|
||||
optional<Beo4Data> decode(RemoteReceiveData src) override;
|
||||
void dump(const Beo4Data &data) override;
|
||||
void encode(RemoteTransmitData *dst, const Beo4Data &data);
|
||||
optional<Beo4Data> decode(RemoteReceiveData src);
|
||||
void dump(const Beo4Data &data);
|
||||
};
|
||||
|
||||
DECLARE_REMOTE_PROTOCOL(Beo4)
|
||||
|
||||
@@ -13,9 +13,9 @@ struct BrennenstuhlData {
|
||||
|
||||
class BrennenstuhlProtocol : public RemoteProtocol<BrennenstuhlData> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const BrennenstuhlData &data) override;
|
||||
optional<BrennenstuhlData> decode(RemoteReceiveData src) override;
|
||||
void dump(const BrennenstuhlData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const BrennenstuhlData &data);
|
||||
optional<BrennenstuhlData> decode(RemoteReceiveData src);
|
||||
void dump(const BrennenstuhlData &data);
|
||||
};
|
||||
|
||||
DECLARE_REMOTE_PROTOCOL(Brennenstuhl)
|
||||
|
||||
@@ -21,9 +21,9 @@ struct ByronSXData {
|
||||
|
||||
class ByronSXProtocol : public RemoteProtocol<ByronSXData> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const ByronSXData &data) override;
|
||||
optional<ByronSXData> decode(RemoteReceiveData src) override;
|
||||
void dump(const ByronSXData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const ByronSXData &data);
|
||||
optional<ByronSXData> decode(RemoteReceiveData src);
|
||||
void dump(const ByronSXData &data);
|
||||
};
|
||||
|
||||
DECLARE_REMOTE_PROTOCOL(ByronSX)
|
||||
|
||||
@@ -19,9 +19,9 @@ struct CanalSatLDData : public CanalSatData {};
|
||||
|
||||
class CanalSatBaseProtocol : public RemoteProtocol<CanalSatData> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const CanalSatData &data) override;
|
||||
optional<CanalSatData> decode(RemoteReceiveData src) override;
|
||||
void dump(const CanalSatData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const CanalSatData &data);
|
||||
optional<CanalSatData> decode(RemoteReceiveData src);
|
||||
void dump(const CanalSatData &data);
|
||||
|
||||
protected:
|
||||
uint16_t frequency_;
|
||||
|
||||
@@ -21,9 +21,9 @@ struct CoolixData {
|
||||
|
||||
class CoolixProtocol : public RemoteProtocol<CoolixData> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const CoolixData &data) override;
|
||||
optional<CoolixData> decode(RemoteReceiveData data) override;
|
||||
void dump(const CoolixData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const CoolixData &data);
|
||||
optional<CoolixData> decode(RemoteReceiveData data);
|
||||
void dump(const CoolixData &data);
|
||||
};
|
||||
|
||||
DECLARE_REMOTE_PROTOCOL(Coolix)
|
||||
|
||||
@@ -13,9 +13,9 @@ struct DishData {
|
||||
|
||||
class DishProtocol : public RemoteProtocol<DishData> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const DishData &data) override;
|
||||
optional<DishData> decode(RemoteReceiveData src) override;
|
||||
void dump(const DishData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const DishData &data);
|
||||
optional<DishData> decode(RemoteReceiveData src);
|
||||
void dump(const DishData &data);
|
||||
};
|
||||
|
||||
DECLARE_REMOTE_PROTOCOL(Dish)
|
||||
|
||||
@@ -20,9 +20,9 @@ struct DooyaData {
|
||||
|
||||
class DooyaProtocol : public RemoteProtocol<DooyaData> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const DooyaData &data) override;
|
||||
optional<DooyaData> decode(RemoteReceiveData src) override;
|
||||
void dump(const DooyaData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const DooyaData &data);
|
||||
optional<DooyaData> decode(RemoteReceiveData src);
|
||||
void dump(const DooyaData &data);
|
||||
};
|
||||
|
||||
DECLARE_REMOTE_PROTOCOL(Dooya)
|
||||
|
||||
@@ -19,9 +19,9 @@ struct DraytonData {
|
||||
|
||||
class DraytonProtocol : public RemoteProtocol<DraytonData> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const DraytonData &data) override;
|
||||
optional<DraytonData> decode(RemoteReceiveData src) override;
|
||||
void dump(const DraytonData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const DraytonData &data);
|
||||
optional<DraytonData> decode(RemoteReceiveData src);
|
||||
void dump(const DraytonData &data);
|
||||
};
|
||||
|
||||
DECLARE_REMOTE_PROTOCOL(Drayton)
|
||||
|
||||
@@ -21,9 +21,9 @@ struct DysonData {
|
||||
|
||||
class DysonProtocol : public RemoteProtocol<DysonData> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const DysonData &data) override;
|
||||
optional<DysonData> decode(RemoteReceiveData src) override;
|
||||
void dump(const DysonData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const DysonData &data);
|
||||
optional<DysonData> decode(RemoteReceiveData src);
|
||||
void dump(const DysonData &data);
|
||||
};
|
||||
|
||||
DECLARE_REMOTE_PROTOCOL(Dyson)
|
||||
|
||||
@@ -31,9 +31,9 @@ class GoboxProtocol : public RemoteProtocol<GoboxData> {
|
||||
void dump_timings_(const RawTimings &timings) const;
|
||||
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const GoboxData &data) override;
|
||||
optional<GoboxData> decode(RemoteReceiveData src) override;
|
||||
void dump(const GoboxData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const GoboxData &data);
|
||||
optional<GoboxData> decode(RemoteReceiveData src);
|
||||
void dump(const GoboxData &data);
|
||||
};
|
||||
|
||||
DECLARE_REMOTE_PROTOCOL(Gobox)
|
||||
|
||||
@@ -13,9 +13,9 @@ struct HaierData {
|
||||
|
||||
class HaierProtocol : public RemoteProtocol<HaierData> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const HaierData &data) override;
|
||||
optional<HaierData> decode(RemoteReceiveData src) override;
|
||||
void dump(const HaierData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const HaierData &data);
|
||||
optional<HaierData> decode(RemoteReceiveData src);
|
||||
void dump(const HaierData &data);
|
||||
|
||||
protected:
|
||||
void encode_byte_(RemoteTransmitData *dst, uint8_t item);
|
||||
|
||||
@@ -14,9 +14,9 @@ struct JVCData {
|
||||
|
||||
class JVCProtocol : public RemoteProtocol<JVCData> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const JVCData &data) override;
|
||||
optional<JVCData> decode(RemoteReceiveData src) override;
|
||||
void dump(const JVCData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const JVCData &data);
|
||||
optional<JVCData> decode(RemoteReceiveData src);
|
||||
void dump(const JVCData &data);
|
||||
};
|
||||
|
||||
DECLARE_REMOTE_PROTOCOL(JVC)
|
||||
|
||||
@@ -24,9 +24,9 @@ struct KeeloqData {
|
||||
|
||||
class KeeloqProtocol : public RemoteProtocol<KeeloqData> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const KeeloqData &data) override;
|
||||
optional<KeeloqData> decode(RemoteReceiveData src) override;
|
||||
void dump(const KeeloqData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const KeeloqData &data);
|
||||
optional<KeeloqData> decode(RemoteReceiveData src);
|
||||
void dump(const KeeloqData &data);
|
||||
};
|
||||
|
||||
DECLARE_REMOTE_PROTOCOL(Keeloq)
|
||||
|
||||
@@ -16,9 +16,9 @@ struct LGData {
|
||||
|
||||
class LGProtocol : public RemoteProtocol<LGData> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const LGData &data) override;
|
||||
optional<LGData> decode(RemoteReceiveData src) override;
|
||||
void dump(const LGData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const LGData &data);
|
||||
optional<LGData> decode(RemoteReceiveData src);
|
||||
void dump(const LGData &data);
|
||||
};
|
||||
|
||||
DECLARE_REMOTE_PROTOCOL(LG)
|
||||
|
||||
@@ -27,9 +27,9 @@ struct MagiQuestData {
|
||||
|
||||
class MagiQuestProtocol : public RemoteProtocol<MagiQuestData> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const MagiQuestData &data) override;
|
||||
optional<MagiQuestData> decode(RemoteReceiveData src) override;
|
||||
void dump(const MagiQuestData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const MagiQuestData &data);
|
||||
optional<MagiQuestData> decode(RemoteReceiveData src);
|
||||
void dump(const MagiQuestData &data);
|
||||
};
|
||||
|
||||
DECLARE_REMOTE_PROTOCOL(MagiQuest)
|
||||
|
||||
@@ -67,9 +67,9 @@ class MideaData {
|
||||
|
||||
class MideaProtocol : public RemoteProtocol<MideaData> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const MideaData &src) override;
|
||||
optional<MideaData> decode(RemoteReceiveData src) override;
|
||||
void dump(const MideaData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const MideaData &src);
|
||||
optional<MideaData> decode(RemoteReceiveData src);
|
||||
void dump(const MideaData &data);
|
||||
};
|
||||
|
||||
DECLARE_REMOTE_PROTOCOL(Midea)
|
||||
|
||||
@@ -13,9 +13,9 @@ struct MirageData {
|
||||
|
||||
class MirageProtocol : public RemoteProtocol<MirageData> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const MirageData &data) override;
|
||||
optional<MirageData> decode(RemoteReceiveData src) override;
|
||||
void dump(const MirageData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const MirageData &data);
|
||||
optional<MirageData> decode(RemoteReceiveData src);
|
||||
void dump(const MirageData &data);
|
||||
|
||||
protected:
|
||||
void encode_byte_(RemoteTransmitData *dst, uint8_t item);
|
||||
|
||||
@@ -14,9 +14,9 @@ struct NECData {
|
||||
|
||||
class NECProtocol : public RemoteProtocol<NECData> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const NECData &data) override;
|
||||
optional<NECData> decode(RemoteReceiveData src) override;
|
||||
void dump(const NECData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const NECData &data);
|
||||
optional<NECData> decode(RemoteReceiveData src);
|
||||
void dump(const NECData &data);
|
||||
};
|
||||
|
||||
DECLARE_REMOTE_PROTOCOL(NEC)
|
||||
|
||||
@@ -24,9 +24,9 @@ class NexaProtocol : public RemoteProtocol<NexaData> {
|
||||
void zero(RemoteTransmitData *dst) const;
|
||||
void sync(RemoteTransmitData *dst) const;
|
||||
|
||||
void encode(RemoteTransmitData *dst, const NexaData &data) override;
|
||||
optional<NexaData> decode(RemoteReceiveData src) override;
|
||||
void dump(const NexaData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const NexaData &data);
|
||||
optional<NexaData> decode(RemoteReceiveData src);
|
||||
void dump(const NexaData &data);
|
||||
};
|
||||
|
||||
DECLARE_REMOTE_PROTOCOL(Nexa)
|
||||
|
||||
@@ -16,9 +16,9 @@ struct PanasonicData {
|
||||
|
||||
class PanasonicProtocol : public RemoteProtocol<PanasonicData> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const PanasonicData &data) override;
|
||||
optional<PanasonicData> decode(RemoteReceiveData src) override;
|
||||
void dump(const PanasonicData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const PanasonicData &data);
|
||||
optional<PanasonicData> decode(RemoteReceiveData src);
|
||||
void dump(const PanasonicData &data);
|
||||
};
|
||||
|
||||
DECLARE_REMOTE_PROTOCOL(Panasonic)
|
||||
|
||||
@@ -13,9 +13,9 @@ struct PioneerData {
|
||||
|
||||
class PioneerProtocol : public RemoteProtocol<PioneerData> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const PioneerData &data) override;
|
||||
optional<PioneerData> decode(RemoteReceiveData src) override;
|
||||
void dump(const PioneerData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const PioneerData &data);
|
||||
optional<PioneerData> decode(RemoteReceiveData src);
|
||||
void dump(const PioneerData &data);
|
||||
};
|
||||
|
||||
DECLARE_REMOTE_PROTOCOL(Pioneer)
|
||||
|
||||
@@ -30,9 +30,9 @@ class ProntoProtocol : public RemoteProtocol<ProntoData> {
|
||||
std::string compensate_and_dump_sequence_(const RawTimings &data, uint16_t timebase);
|
||||
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const ProntoData &data) override;
|
||||
optional<ProntoData> decode(RemoteReceiveData src) override;
|
||||
void dump(const ProntoData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const ProntoData &data);
|
||||
optional<ProntoData> decode(RemoteReceiveData src);
|
||||
void dump(const ProntoData &data);
|
||||
};
|
||||
|
||||
DECLARE_REMOTE_PROTOCOL(Pronto)
|
||||
|
||||
@@ -14,9 +14,9 @@ struct RC5Data {
|
||||
|
||||
class RC5Protocol : public RemoteProtocol<RC5Data> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const RC5Data &data) override;
|
||||
optional<RC5Data> decode(RemoteReceiveData src) override;
|
||||
void dump(const RC5Data &data) override;
|
||||
void encode(RemoteTransmitData *dst, const RC5Data &data);
|
||||
optional<RC5Data> decode(RemoteReceiveData src);
|
||||
void dump(const RC5Data &data);
|
||||
};
|
||||
|
||||
DECLARE_REMOTE_PROTOCOL(RC5)
|
||||
|
||||
@@ -15,9 +15,9 @@ struct RC6Data {
|
||||
|
||||
class RC6Protocol : public RemoteProtocol<RC6Data> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const RC6Data &data) override;
|
||||
optional<RC6Data> decode(RemoteReceiveData src) override;
|
||||
void dump(const RC6Data &data) override;
|
||||
void encode(RemoteTransmitData *dst, const RC6Data &data);
|
||||
optional<RC6Data> decode(RemoteReceiveData src);
|
||||
void dump(const RC6Data &data);
|
||||
};
|
||||
|
||||
DECLARE_REMOTE_PROTOCOL(RC6)
|
||||
|
||||
@@ -1,30 +1,12 @@
|
||||
#include "rc_switch_protocol.h"
|
||||
|
||||
#include <iterator>
|
||||
#include "esphome/core/log.h"
|
||||
|
||||
namespace esphome::remote_base {
|
||||
|
||||
static const char *const TAG = "remote.rc_switch";
|
||||
|
||||
const RCSwitchBase RC_SWITCH_PROTOCOLS[9] = {RCSwitchBase(0, 0, 0, 0, 0, 0, false),
|
||||
RCSwitchBase(350, 10850, 350, 1050, 1050, 350, false),
|
||||
RCSwitchBase(650, 6500, 650, 1300, 1300, 650, false),
|
||||
RCSwitchBase(3000, 7100, 400, 1100, 900, 600, false),
|
||||
RCSwitchBase(380, 2280, 380, 1140, 1140, 380, false),
|
||||
RCSwitchBase(3000, 7000, 500, 1000, 1000, 500, false),
|
||||
RCSwitchBase(10350, 450, 450, 900, 900, 450, true),
|
||||
RCSwitchBase(300, 9300, 150, 900, 900, 150, false),
|
||||
RCSwitchBase(250, 2500, 250, 1250, 250, 250, false)};
|
||||
|
||||
RCSwitchBase::RCSwitchBase(uint32_t sync_high, uint32_t sync_low, uint32_t zero_high, uint32_t zero_low,
|
||||
uint32_t one_high, uint32_t one_low, bool inverted)
|
||||
: sync_high_(sync_high),
|
||||
sync_low_(sync_low),
|
||||
zero_high_(zero_high),
|
||||
zero_low_(zero_low),
|
||||
one_high_(one_high),
|
||||
one_low_(one_low),
|
||||
inverted_(inverted) {}
|
||||
|
||||
void RCSwitchBase::one(RemoteTransmitData *dst) const {
|
||||
if (!this->inverted_) {
|
||||
dst->mark(this->one_high_);
|
||||
@@ -133,11 +115,11 @@ bool RCSwitchBase::decode(RemoteReceiveData &src, uint64_t *out_data, uint8_t *o
|
||||
optional<RCSwitchData> RCSwitchBase::decode(RemoteReceiveData &src) const {
|
||||
RCSwitchData out;
|
||||
uint8_t out_nbits;
|
||||
for (uint8_t i = 1; i <= 8; i++) {
|
||||
for (size_t i = 1; i < std::size(RC_SWITCH_PROTOCOLS); i++) {
|
||||
src.reset();
|
||||
const RCSwitchBase *protocol = &RC_SWITCH_PROTOCOLS[i];
|
||||
if (protocol->decode(src, &out.code, &out_nbits) && out_nbits >= 3) {
|
||||
out.protocol = i;
|
||||
out.protocol = static_cast<uint8_t>(i);
|
||||
return out;
|
||||
}
|
||||
}
|
||||
@@ -246,7 +228,7 @@ bool RCSwitchRawReceiver::matches(RemoteReceiveData src) {
|
||||
return decoded_nbits == this->nbits_ && (decoded_code & this->mask_) == (this->code_ & this->mask_);
|
||||
}
|
||||
bool RCSwitchDumper::dump(RemoteReceiveData src) {
|
||||
for (uint8_t i = 1; i <= 8; i++) {
|
||||
for (size_t i = 1; i < std::size(RC_SWITCH_PROTOCOLS); i++) {
|
||||
src.reset();
|
||||
uint64_t out_data;
|
||||
uint8_t out_nbits;
|
||||
@@ -257,7 +239,7 @@ bool RCSwitchDumper::dump(RemoteReceiveData src) {
|
||||
buffer[j] = (out_data & ((uint64_t) 1 << (out_nbits - j - 1))) ? '1' : '0';
|
||||
|
||||
buffer[out_nbits] = '\0';
|
||||
ESP_LOGI(TAG, "Received RCSwitch Raw: protocol=%u data='%s'", i, buffer);
|
||||
ESP_LOGI(TAG, "Received RCSwitch Raw: protocol=%u data='%s'", static_cast<unsigned>(i), buffer);
|
||||
|
||||
// only send first decoded protocol
|
||||
return true;
|
||||
|
||||
@@ -16,9 +16,16 @@ class RCSwitchBase {
|
||||
public:
|
||||
using ProtocolData = RCSwitchData;
|
||||
|
||||
RCSwitchBase() = default;
|
||||
RCSwitchBase(uint32_t sync_high, uint32_t sync_low, uint32_t zero_high, uint32_t zero_low, uint32_t one_high,
|
||||
uint32_t one_low, bool inverted);
|
||||
constexpr RCSwitchBase() = default;
|
||||
constexpr RCSwitchBase(uint32_t sync_high, uint32_t sync_low, uint32_t zero_high, uint32_t zero_low,
|
||||
uint32_t one_high, uint32_t one_low, bool inverted)
|
||||
: sync_high_(sync_high),
|
||||
sync_low_(sync_low),
|
||||
zero_high_(zero_high),
|
||||
zero_low_(zero_low),
|
||||
one_high_(one_high),
|
||||
one_low_(one_low),
|
||||
inverted_(inverted) {}
|
||||
|
||||
void one(RemoteTransmitData *dst) const;
|
||||
|
||||
@@ -58,10 +65,21 @@ class RCSwitchBase {
|
||||
uint32_t zero_low_{};
|
||||
uint32_t one_high_{};
|
||||
uint32_t one_low_{};
|
||||
bool inverted_{};
|
||||
uint32_t inverted_{}; // bool widened so every field is a word: the table is read from flash
|
||||
};
|
||||
|
||||
extern const RCSwitchBase RC_SWITCH_PROTOCOLS[9];
|
||||
// Constant-initialized and kept in flash on every platform; all fields are 32-bit so ESP8266 can read it in place
|
||||
inline constexpr RCSwitchBase RC_SWITCH_PROTOCOLS[] PROGMEM = {
|
||||
{0, 0, 0, 0, 0, 0, false},
|
||||
{350, 10850, 350, 1050, 1050, 350, false},
|
||||
{650, 6500, 650, 1300, 1300, 650, false},
|
||||
{3000, 7100, 400, 1100, 900, 600, false},
|
||||
{380, 2280, 380, 1140, 1140, 380, false},
|
||||
{3000, 7000, 500, 1000, 1000, 500, false},
|
||||
{10350, 450, 450, 900, 900, 450, true},
|
||||
{300, 9300, 150, 900, 900, 150, false},
|
||||
{250, 2500, 250, 1250, 250, 250, false},
|
||||
};
|
||||
|
||||
uint64_t decode_binary_string(const std::string &data);
|
||||
|
||||
|
||||
@@ -99,29 +99,47 @@ bool RemoteReceiverBinarySensorBase::on_receive(RemoteReceiveData src) {
|
||||
|
||||
/* RemoteReceiverBase */
|
||||
|
||||
// Slots are counted at code generation; a registration from C++ setup() has none
|
||||
#ifdef REMOTE_BASE_LISTENER_COUNT
|
||||
void RemoteReceiverBase::register_listener(RemoteReceiverListener *listener) {
|
||||
if (this->listeners_.size() == REMOTE_BASE_LISTENER_COUNT) {
|
||||
ESP_LOGE(TAG, "No %s slot: register it from to_code() with remote_base.add_%s", LOG_STR_LITERAL("listener"),
|
||||
LOG_STR_LITERAL("listener"));
|
||||
return;
|
||||
}
|
||||
this->listeners_.push_back(listener);
|
||||
}
|
||||
#endif
|
||||
|
||||
#ifdef REMOTE_BASE_DUMPER_COUNT
|
||||
void RemoteReceiverBase::register_dumper(RemoteReceiverDumperBase *dumper) {
|
||||
if (dumper->is_secondary()) {
|
||||
this->secondary_dumpers_.push_back(dumper);
|
||||
} else {
|
||||
this->dumpers_.push_back(dumper);
|
||||
this->secondary_dumper_ = dumper;
|
||||
return;
|
||||
}
|
||||
if (this->dumpers_.size() == REMOTE_BASE_DUMPER_COUNT) {
|
||||
ESP_LOGE(TAG, "No %s slot: register it from to_code() with remote_base.add_%s", LOG_STR_LITERAL("dumper"),
|
||||
LOG_STR_LITERAL("dumper"));
|
||||
return;
|
||||
}
|
||||
this->dumpers_.push_back(dumper);
|
||||
}
|
||||
#endif
|
||||
|
||||
void RemoteReceiverBase::call_listeners_() {
|
||||
void RemoteReceiverBase::call_listeners_dumpers_() {
|
||||
#ifdef REMOTE_BASE_LISTENER_COUNT
|
||||
for (auto *listener : this->listeners_)
|
||||
listener->on_receive(RemoteReceiveData(this->temp_, this->tolerance_, this->tolerance_mode_));
|
||||
}
|
||||
|
||||
void RemoteReceiverBase::call_dumpers_() {
|
||||
#endif
|
||||
#ifdef REMOTE_BASE_DUMPER_COUNT
|
||||
bool success = false;
|
||||
for (auto *dumper : this->dumpers_) {
|
||||
if (dumper->dump(RemoteReceiveData(this->temp_, this->tolerance_, this->tolerance_mode_)))
|
||||
success = true;
|
||||
}
|
||||
if (!success) {
|
||||
for (auto *dumper : this->secondary_dumpers_)
|
||||
dumper->dump(RemoteReceiveData(this->temp_, this->tolerance_, this->tolerance_mode_));
|
||||
}
|
||||
if (!success && this->secondary_dumper_ != nullptr)
|
||||
this->secondary_dumper_->dump(RemoteReceiveData(this->temp_, this->tolerance_, this->tolerance_mode_));
|
||||
#endif
|
||||
}
|
||||
|
||||
void RemoteReceiverBinarySensorBase::dump_config() { LOG_BINARY_SENSOR("", "Remote Receiver Binary Sensor", this); }
|
||||
|
||||
@@ -1,12 +1,14 @@
|
||||
#pragma once
|
||||
|
||||
#include <concepts>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "esphome/components/binary_sensor/binary_sensor.h"
|
||||
#include "esphome/core/automation.h"
|
||||
#include "esphome/core/component.h"
|
||||
#include "esphome/core/hal.h"
|
||||
#include "esphome/core/helpers.h"
|
||||
|
||||
namespace esphome::remote_base {
|
||||
|
||||
@@ -141,6 +143,22 @@ class RemoteRMTChannel {
|
||||
#endif // SOC_RMT_SUPPORTED
|
||||
#endif // USE_ESP32
|
||||
|
||||
// Protocol shapes, checked where a protocol is used so a missing method fails at the use site
|
||||
// instead of deep inside a template body. Receive-only protocols such as RCSwitchBase decode
|
||||
// without encoding.
|
||||
template<typename T>
|
||||
concept RemoteProtocolDecoder = requires(T proto, RemoteReceiveData src) {
|
||||
{ proto.decode(src) } -> std::same_as<optional<typename T::ProtocolData>>;
|
||||
};
|
||||
template<typename T>
|
||||
concept RemoteProtocolDumper = RemoteProtocolDecoder<T> && requires(T proto, const typename T::ProtocolData &data) {
|
||||
proto.dump(data);
|
||||
};
|
||||
template<typename T>
|
||||
concept RemoteProtocolEncoder = requires(T proto, RemoteTransmitData *dst, const typename T::ProtocolData &data) {
|
||||
proto.encode(dst, data);
|
||||
};
|
||||
|
||||
class RemoteTransmitterBase : public RemoteComponentBase {
|
||||
public:
|
||||
RemoteTransmitterBase(InternalGPIOPin *pin) : RemoteComponentBase(pin) {}
|
||||
@@ -162,7 +180,7 @@ class RemoteTransmitterBase : public RemoteComponentBase {
|
||||
this->temp_.reset();
|
||||
return TransmitCall(this);
|
||||
}
|
||||
template<typename Protocol>
|
||||
template<RemoteProtocolEncoder Protocol>
|
||||
void transmit(const Protocol::ProtocolData &data, uint32_t send_times = 1, uint32_t send_wait = 0) {
|
||||
auto call = this->transmit();
|
||||
Protocol().encode(call.get_data(), data);
|
||||
@@ -194,24 +212,37 @@ class RemoteReceiverDumperBase {
|
||||
class RemoteReceiverBase : public RemoteComponentBase {
|
||||
public:
|
||||
RemoteReceiverBase(InternalGPIOPin *pin) : RemoteComponentBase(pin) {}
|
||||
void register_listener(RemoteReceiverListener *listener) { this->listeners_.push_back(listener); }
|
||||
// Slots are counted at code generation; without one the call fails at compile time with the same message
|
||||
// the runtime check logs
|
||||
#ifdef REMOTE_BASE_LISTENER_COUNT
|
||||
void register_listener(RemoteReceiverListener *listener);
|
||||
#else
|
||||
template<typename T> void register_listener(T *) {
|
||||
static_assert(sizeof(T) == 0, "No listener slot: register it from to_code() with remote_base.add_listener");
|
||||
}
|
||||
#endif
|
||||
#ifdef REMOTE_BASE_DUMPER_COUNT
|
||||
void register_dumper(RemoteReceiverDumperBase *dumper);
|
||||
#else
|
||||
template<typename T> void register_dumper(T *) {
|
||||
static_assert(sizeof(T) == 0, "No dumper slot: register it from to_code() with remote_base.add_dumper");
|
||||
}
|
||||
#endif
|
||||
void set_tolerance(uint32_t tolerance, ToleranceMode tolerance_mode) {
|
||||
this->tolerance_ = tolerance;
|
||||
this->tolerance_mode_ = tolerance_mode;
|
||||
}
|
||||
|
||||
protected:
|
||||
void call_listeners_();
|
||||
void call_dumpers_();
|
||||
void call_listeners_dumpers_() {
|
||||
this->call_listeners_();
|
||||
this->call_dumpers_();
|
||||
}
|
||||
void call_listeners_dumpers_();
|
||||
|
||||
std::vector<RemoteReceiverListener *> listeners_;
|
||||
std::vector<RemoteReceiverDumperBase *> dumpers_;
|
||||
std::vector<RemoteReceiverDumperBase *> secondary_dumpers_;
|
||||
#ifdef REMOTE_BASE_LISTENER_COUNT
|
||||
StaticVector<RemoteReceiverListener *, REMOTE_BASE_LISTENER_COUNT> listeners_;
|
||||
#endif
|
||||
#ifdef REMOTE_BASE_DUMPER_COUNT
|
||||
StaticVector<RemoteReceiverDumperBase *, REMOTE_BASE_DUMPER_COUNT> dumpers_;
|
||||
RemoteReceiverDumperBase *secondary_dumper_{nullptr}; // runs only when no primary dumper matched
|
||||
#endif
|
||||
RawTimings temp_;
|
||||
uint32_t tolerance_{25};
|
||||
ToleranceMode tolerance_mode_{TOLERANCE_MODE_PERCENTAGE};
|
||||
@@ -229,15 +260,14 @@ class RemoteReceiverBinarySensorBase : public binary_sensor::BinarySensorInitial
|
||||
|
||||
/* TEMPLATES */
|
||||
|
||||
// Protocols are used only through their concrete type (see the RemoteProtocol* concepts); encode/decode/dump
|
||||
// stay non-virtual so unused ones link out
|
||||
template<typename T> class RemoteProtocol {
|
||||
public:
|
||||
using ProtocolData = T;
|
||||
virtual void encode(RemoteTransmitData *dst, const ProtocolData &data) = 0;
|
||||
virtual optional<ProtocolData> decode(RemoteReceiveData src) = 0;
|
||||
virtual void dump(const ProtocolData &data) = 0;
|
||||
};
|
||||
|
||||
template<typename T> class RemoteReceiverBinarySensor : public RemoteReceiverBinarySensorBase {
|
||||
template<RemoteProtocolDecoder T> class RemoteReceiverBinarySensor : public RemoteReceiverBinarySensorBase {
|
||||
public:
|
||||
RemoteReceiverBinarySensor() : RemoteReceiverBinarySensorBase() {}
|
||||
|
||||
@@ -255,7 +285,7 @@ template<typename T> class RemoteReceiverBinarySensor : public RemoteReceiverBin
|
||||
T::ProtocolData data_;
|
||||
};
|
||||
|
||||
template<typename T>
|
||||
template<RemoteProtocolDecoder T>
|
||||
class RemoteReceiverTrigger final : public Trigger<typename T::ProtocolData>, public RemoteReceiverListener {
|
||||
protected:
|
||||
bool on_receive(RemoteReceiveData src) override {
|
||||
@@ -276,7 +306,7 @@ class RemoteTransmittable {
|
||||
void set_transmitter(RemoteTransmitterBase *transmitter) { this->transmitter_ = transmitter; }
|
||||
|
||||
protected:
|
||||
template<typename Protocol>
|
||||
template<RemoteProtocolEncoder Protocol>
|
||||
void transmit_(const Protocol::ProtocolData &data, uint32_t send_times = 1, uint32_t send_wait = 0) {
|
||||
this->transmitter_->transmit<Protocol>(data, send_times, send_wait);
|
||||
}
|
||||
@@ -298,7 +328,7 @@ template<typename... Ts> class RemoteTransmitterActionBase : public RemoteTransm
|
||||
virtual void encode(RemoteTransmitData *dst, Ts... x) = 0;
|
||||
};
|
||||
|
||||
template<typename T> class RemoteReceiverDumper : public RemoteReceiverDumperBase {
|
||||
template<RemoteProtocolDumper T> class RemoteReceiverDumper : public RemoteReceiverDumperBase {
|
||||
public:
|
||||
bool dump(RemoteReceiveData src) override {
|
||||
auto proto = T();
|
||||
|
||||
@@ -12,9 +12,9 @@ struct RoombaData {
|
||||
|
||||
class RoombaProtocol : public RemoteProtocol<RoombaData> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const RoombaData &data) override;
|
||||
optional<RoombaData> decode(RemoteReceiveData src) override;
|
||||
void dump(const RoombaData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const RoombaData &data);
|
||||
optional<RoombaData> decode(RemoteReceiveData src);
|
||||
void dump(const RoombaData &data);
|
||||
};
|
||||
|
||||
DECLARE_REMOTE_PROTOCOL(Roomba)
|
||||
|
||||
@@ -16,9 +16,9 @@ struct Samsung36Data {
|
||||
|
||||
class Samsung36Protocol : public RemoteProtocol<Samsung36Data> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const Samsung36Data &data) override;
|
||||
optional<Samsung36Data> decode(RemoteReceiveData src) override;
|
||||
void dump(const Samsung36Data &data) override;
|
||||
void encode(RemoteTransmitData *dst, const Samsung36Data &data);
|
||||
optional<Samsung36Data> decode(RemoteReceiveData src);
|
||||
void dump(const Samsung36Data &data);
|
||||
};
|
||||
|
||||
DECLARE_REMOTE_PROTOCOL(Samsung36)
|
||||
|
||||
@@ -14,9 +14,9 @@ struct SamsungData {
|
||||
|
||||
class SamsungProtocol : public RemoteProtocol<SamsungData> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const SamsungData &data) override;
|
||||
optional<SamsungData> decode(RemoteReceiveData src) override;
|
||||
void dump(const SamsungData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const SamsungData &data);
|
||||
optional<SamsungData> decode(RemoteReceiveData src);
|
||||
void dump(const SamsungData &data);
|
||||
};
|
||||
|
||||
DECLARE_REMOTE_PROTOCOL(Samsung)
|
||||
|
||||
@@ -16,9 +16,9 @@ struct SonyData {
|
||||
|
||||
class SonyProtocol : public RemoteProtocol<SonyData> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const SonyData &data) override;
|
||||
optional<SonyData> decode(RemoteReceiveData src) override;
|
||||
void dump(const SonyData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const SonyData &data);
|
||||
optional<SonyData> decode(RemoteReceiveData src);
|
||||
void dump(const SonyData &data);
|
||||
};
|
||||
|
||||
DECLARE_REMOTE_PROTOCOL(Sony)
|
||||
|
||||
@@ -17,9 +17,9 @@ struct SymphonyData {
|
||||
|
||||
class SymphonyProtocol : public RemoteProtocol<SymphonyData> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const SymphonyData &data) override;
|
||||
optional<SymphonyData> decode(RemoteReceiveData src) override;
|
||||
void dump(const SymphonyData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const SymphonyData &data);
|
||||
optional<SymphonyData> decode(RemoteReceiveData src);
|
||||
void dump(const SymphonyData &data);
|
||||
};
|
||||
|
||||
DECLARE_REMOTE_PROTOCOL(Symphony)
|
||||
|
||||
@@ -14,9 +14,9 @@ struct ToshibaAcData {
|
||||
|
||||
class ToshibaAcProtocol : public RemoteProtocol<ToshibaAcData> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const ToshibaAcData &data) override;
|
||||
optional<ToshibaAcData> decode(RemoteReceiveData src) override;
|
||||
void dump(const ToshibaAcData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const ToshibaAcData &data);
|
||||
optional<ToshibaAcData> decode(RemoteReceiveData src);
|
||||
void dump(const ToshibaAcData &data);
|
||||
};
|
||||
|
||||
DECLARE_REMOTE_PROTOCOL(ToshibaAc)
|
||||
|
||||
@@ -16,9 +16,9 @@ struct TotoData {
|
||||
|
||||
class TotoProtocol : public RemoteProtocol<TotoData> {
|
||||
public:
|
||||
void encode(RemoteTransmitData *dst, const TotoData &data) override;
|
||||
optional<TotoData> decode(RemoteReceiveData src) override;
|
||||
void dump(const TotoData &data) override;
|
||||
void encode(RemoteTransmitData *dst, const TotoData &data);
|
||||
optional<TotoData> decode(RemoteReceiveData src);
|
||||
void dump(const TotoData &data);
|
||||
};
|
||||
|
||||
DECLARE_REMOTE_PROTOCOL(Toto)
|
||||
|
||||
@@ -221,11 +221,11 @@ async def to_code(config: ConfigType) -> None:
|
||||
|
||||
dumpers = await remote_base.build_dumpers(config[CONF_DUMP])
|
||||
for dumper in dumpers:
|
||||
cg.add(var.register_dumper(dumper))
|
||||
remote_base.add_dumper(var, dumper)
|
||||
|
||||
triggers = await remote_base.build_triggers(config)
|
||||
for trigger in triggers:
|
||||
cg.add(var.register_listener(trigger))
|
||||
remote_base.add_listener(var, trigger)
|
||||
await cg.register_component(var, config)
|
||||
|
||||
cg.add(
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import esphome.codegen as cg
|
||||
from esphome.components import climate_ir
|
||||
from esphome.components import climate_ir, remote_base
|
||||
import esphome.config_validation as cv
|
||||
from esphome.const import CONF_MODEL
|
||||
from esphome.types import ConfigType
|
||||
@@ -26,5 +26,6 @@ CONFIG_SCHEMA = climate_ir.climate_ir_with_receiver_schema(ToshibaClimate).exten
|
||||
|
||||
|
||||
async def to_code(config: ConfigType) -> None:
|
||||
remote_base.request_protocol("toshiba_ac") # used from C++
|
||||
var = await climate_ir.new_climate_ir(config)
|
||||
cg.add(var.set_model(config[CONF_MODEL]))
|
||||
|
||||
@@ -137,6 +137,43 @@
|
||||
#define MICRONOVA_LISTENER_COUNT 1
|
||||
#define USE_MICRONOVA_WRITER
|
||||
#define MK2PVROUTER_LISTENER_COUNT 1
|
||||
#define REMOTE_BASE_DUMPER_COUNT 1
|
||||
#define REMOTE_BASE_LISTENER_COUNT 1
|
||||
#define USE_REMOTE_PROTOCOL_ABBWELCOME
|
||||
#define USE_REMOTE_PROTOCOL_AEHA
|
||||
#define USE_REMOTE_PROTOCOL_BEO4
|
||||
#define USE_REMOTE_PROTOCOL_BRENNENSTUHL
|
||||
#define USE_REMOTE_PROTOCOL_BYRONSX
|
||||
#define USE_REMOTE_PROTOCOL_CANALSAT
|
||||
#define USE_REMOTE_PROTOCOL_COOLIX
|
||||
#define USE_REMOTE_PROTOCOL_DISH
|
||||
#define USE_REMOTE_PROTOCOL_DOOYA
|
||||
#define USE_REMOTE_PROTOCOL_DRAYTON
|
||||
#define USE_REMOTE_PROTOCOL_DYSON
|
||||
#define USE_REMOTE_PROTOCOL_GOBOX
|
||||
#define USE_REMOTE_PROTOCOL_HAIER
|
||||
#define USE_REMOTE_PROTOCOL_JVC
|
||||
#define USE_REMOTE_PROTOCOL_KEELOQ
|
||||
#define USE_REMOTE_PROTOCOL_LG
|
||||
#define USE_REMOTE_PROTOCOL_MAGIQUEST
|
||||
#define USE_REMOTE_PROTOCOL_MIDEA
|
||||
#define USE_REMOTE_PROTOCOL_MIRAGE
|
||||
#define USE_REMOTE_PROTOCOL_NEC
|
||||
#define USE_REMOTE_PROTOCOL_NEXA
|
||||
#define USE_REMOTE_PROTOCOL_PANASONIC
|
||||
#define USE_REMOTE_PROTOCOL_PIONEER
|
||||
#define USE_REMOTE_PROTOCOL_PRONTO
|
||||
#define USE_REMOTE_PROTOCOL_RAW
|
||||
#define USE_REMOTE_PROTOCOL_RC5
|
||||
#define USE_REMOTE_PROTOCOL_RC6
|
||||
#define USE_REMOTE_PROTOCOL_RC_SWITCH
|
||||
#define USE_REMOTE_PROTOCOL_ROOMBA
|
||||
#define USE_REMOTE_PROTOCOL_SAMSUNG
|
||||
#define USE_REMOTE_PROTOCOL_SAMSUNG36
|
||||
#define USE_REMOTE_PROTOCOL_SONY
|
||||
#define USE_REMOTE_PROTOCOL_SYMPHONY
|
||||
#define USE_REMOTE_PROTOCOL_TOSHIBA_AC
|
||||
#define USE_REMOTE_PROTOCOL_TOTO
|
||||
#define SERIAL_PROXY_COUNT 2
|
||||
#define SNTP_SERVER_COUNT 3
|
||||
#define USE_MEDIA_PLAYER
|
||||
|
||||
+294
-215
@@ -28,11 +28,6 @@ class WireType(IntEnum):
|
||||
END_GROUP = 4 # groups (deprecated)
|
||||
FIXED32 = 5 # fixed32, sfixed32, float
|
||||
|
||||
@property
|
||||
def cpp_name(self) -> str:
|
||||
"""The matching constant in proto.h."""
|
||||
return f"WIRE_TYPE_{self.name}"
|
||||
|
||||
|
||||
# Generate with
|
||||
# protoc --python_out=script/api_protobuf -I esphome/components/api/ api_options.proto
|
||||
@@ -131,10 +126,9 @@ def camel_to_snake(name: str) -> str:
|
||||
return re.sub("([a-z0-9])([A-Z])", r"\1_\2", s1).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))});"
|
||||
def force_str(force: bool) -> str:
|
||||
"""Convert a boolean force value to string format for C++ code."""
|
||||
return str(force).lower()
|
||||
|
||||
|
||||
class TypeInfo(ABC):
|
||||
@@ -229,39 +223,55 @@ class TypeInfo(ABC):
|
||||
def class_member(self) -> str:
|
||||
return f"{self.cpp_type} {self.field_name}{{{self.default_value}}};"
|
||||
|
||||
def decode_case(self, body: str) -> str:
|
||||
"""Emit one decode_field() case, keyed on the field's wire tag."""
|
||||
return f"case proto_tag({self.number}, {self.wire_type.cpp_name}):\n" + indent(
|
||||
f"{body}\nbreak;"
|
||||
)
|
||||
@property
|
||||
def decode_varint_content(self) -> str:
|
||||
content = self.decode_varint
|
||||
if content is None:
|
||||
return None
|
||||
return f"case {self.number}: this->{self.field_name} = {content}; break;"
|
||||
|
||||
# Expression that reads this field from `value`; None when the type is never decoded.
|
||||
decode_expr: str | None = None
|
||||
|
||||
def _decode_store(self, expr: str) -> str:
|
||||
return f"this->{self.field_name} = {expr};"
|
||||
decode_varint = None
|
||||
|
||||
@property
|
||||
def decode_content(self) -> str | None:
|
||||
"""The decode_field() case for this field, or None when it is never decoded."""
|
||||
expr = self.decode_expr
|
||||
return None if expr is None else self.decode_case(self._decode_store(expr))
|
||||
def decode_length_content(self) -> str:
|
||||
content = self.decode_length
|
||||
if content is None:
|
||||
return None
|
||||
return f"case {self.number}: this->{self.field_name} = {content}; break;"
|
||||
|
||||
decode_length = None
|
||||
|
||||
@property
|
||||
def decode_32bit_content(self) -> str:
|
||||
content = self.decode_32bit
|
||||
if content is None:
|
||||
return None
|
||||
return f"case {self.number}: this->{self.field_name} = {content}; break;"
|
||||
|
||||
decode_32bit = None
|
||||
|
||||
@property
|
||||
def decode_64bit_content(self) -> str:
|
||||
content = self.decode_64bit
|
||||
if content is None:
|
||||
return None
|
||||
return f"case {self.number}: this->{self.field_name} = {content}; break;"
|
||||
|
||||
decode_64bit = None
|
||||
|
||||
# Mapping from encode_func to raw encode expression template.
|
||||
# When a forced field has a single-byte tag, the code generator emits
|
||||
# write_raw_byte(tag) + raw encode instead of the full encode_* method,
|
||||
# eliminating the zero-check branch and encode_field_raw indirection.
|
||||
# {value} is replaced with the actual field expression.
|
||||
RAW_ENCODE_MAP: dict[str, tuple[str, str]] = {
|
||||
"encode_uint32": ("encode_varint_raw", "{value}"),
|
||||
"encode_uint64": ("encode_varint_raw_64", "{value}"),
|
||||
"encode_sint32": ("encode_varint_raw_short", "encode_zigzag32({value})"),
|
||||
"encode_sint64": ("encode_varint_raw_64", "encode_zigzag64({value})"),
|
||||
"encode_int64": ("encode_varint_raw_64", "static_cast<uint64_t>({value})"),
|
||||
"encode_bool": ("write_raw_byte", "{value} ? 0x01 : 0x00"),
|
||||
RAW_ENCODE_MAP: dict[str, str] = {
|
||||
"encode_uint32": "ProtoEncode::encode_varint_raw(pos, {value});",
|
||||
"encode_uint64": "ProtoEncode::encode_varint_raw_64(pos, {value});",
|
||||
"encode_sint32": "ProtoEncode::encode_varint_raw_short(pos, encode_zigzag32({value}));",
|
||||
"encode_sint64": "ProtoEncode::encode_varint_raw_64(pos, encode_zigzag64({value}));",
|
||||
"encode_int64": "ProtoEncode::encode_varint_raw_64(pos, static_cast<uint64_t>({value}));",
|
||||
"encode_bool": "ProtoEncode::write_raw_byte(pos, {value} ? 0x01 : 0x00);",
|
||||
}
|
||||
# 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:
|
||||
"""Try to emit a precomputed-tag encode for a field.
|
||||
@@ -278,17 +288,12 @@ class TypeInfo(ABC):
|
||||
return None
|
||||
max_val = self.max_value
|
||||
# Only use RAW_ENCODE_MAP for forced fields or fields with max_value
|
||||
raw = None
|
||||
raw_expr = None
|
||||
if self.force or max_val is not None:
|
||||
raw = self.RAW_ENCODE_MAP.get(self.encode_func)
|
||||
if raw is None:
|
||||
raw_expr = self.RAW_ENCODE_MAP.get(self.encode_func)
|
||||
if raw_expr is None:
|
||||
return None
|
||||
func, arg = raw
|
||||
body = (
|
||||
_encode_call("write_raw_byte", str(tag))
|
||||
+ "\n"
|
||||
+ _encode_call(func, arg.format(value=value_expr))
|
||||
)
|
||||
body = f"ProtoEncode::write_raw_byte(pos, {tag});\n{raw_expr.format(value=value_expr)}"
|
||||
if self.force:
|
||||
return body
|
||||
# Non-forced with max_value: inline zero-check + raw encode
|
||||
@@ -309,44 +314,23 @@ class TypeInfo(ABC):
|
||||
return None
|
||||
# When max_len < 128, length varint is always 1 byte
|
||||
len_encode = (
|
||||
_encode_call("write_raw_byte", f"static_cast<uint8_t>({len_expr})")
|
||||
f"ProtoEncode::write_raw_byte(pos, static_cast<uint8_t>({len_expr}));"
|
||||
if max_len is not None and max_len < 128
|
||||
else _encode_call("encode_varint_raw", len_expr)
|
||||
else f"ProtoEncode::encode_varint_raw(pos, {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: str) -> str | None:
|
||||
"""Single-byte tag fixed32 write, or None for other types and multi-byte tags."""
|
||||
tag = self.calculate_tag()
|
||||
if self.fixed32_value_template is None or tag >= 128:
|
||||
return None
|
||||
value_expr = self.fixed32_value_template.format(value=value)
|
||||
if self.force:
|
||||
return _encode_call("write_tag_and_fixed32", str(tag), value_expr)
|
||||
return (
|
||||
f"if (uint32_t raw = {value_expr}; raw != 0) [[likely]] {{\n"
|
||||
f" {_encode_call('write_tag_and_fixed32', str(tag), 'raw')}\n"
|
||||
"}"
|
||||
f"ProtoEncode::write_raw_byte(pos, {tag});\n"
|
||||
f"{len_encode}\n"
|
||||
f"ProtoEncode::encode_raw(pos, {data_expr}, {len_expr});"
|
||||
)
|
||||
|
||||
@property
|
||||
def encode_content(self) -> str:
|
||||
value = f"this->{self.field_name}"
|
||||
if result := self._encode_with_precomputed_tag(value):
|
||||
if result := self._encode_with_precomputed_tag(f"this->{self.field_name}"):
|
||||
return result
|
||||
if result := self._encode_fixed32_with_precomputed_tag(value):
|
||||
return result
|
||||
return _encode_call(self.encode_func, str(self.number), value, force=self.force)
|
||||
|
||||
def encode_element(self, number: int, element: str) -> str:
|
||||
"""Encode one element of a repeated field; elements are always written."""
|
||||
return _encode_call(self.encode_func, str(number), element, force=True)
|
||||
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});"
|
||||
|
||||
encode_func = None
|
||||
|
||||
@@ -621,6 +605,7 @@ class DoubleType(FixedSizeTypeMixin, TypeInfo):
|
||||
# Unsupported but defined for completeness
|
||||
cpp_type = "double"
|
||||
default_value = "0.0"
|
||||
decode_64bit = "value.as_double()"
|
||||
encode_func = "encode_double"
|
||||
wire_type = WireType.FIXED64 # Uses wire type 1 according to protobuf spec
|
||||
|
||||
@@ -646,12 +631,10 @@ class DoubleType(FixedSizeTypeMixin, TypeInfo):
|
||||
class FloatType(FixedSizeTypeMixin, TypeInfo):
|
||||
cpp_type = "float"
|
||||
default_value = "0.0f"
|
||||
decode_expr = "value.as_float()"
|
||||
decode_32bit = "value.as_float()"
|
||||
encode_func = "encode_float"
|
||||
wire_type = WireType.FIXED32 # Uses wire type 5
|
||||
|
||||
fixed32_value_template = "float_to_raw({value})"
|
||||
|
||||
def dump(self, name: str) -> str:
|
||||
o = f'snprintf(buffer, sizeof(buffer), "%g", {name});\n'
|
||||
o += "out.append(buffer);"
|
||||
@@ -675,7 +658,7 @@ class Int64Type(VarintTypeMixin, TypeInfo):
|
||||
cpp_type = "int64_t"
|
||||
_varint_max_bits = 64
|
||||
default_value = "0"
|
||||
decode_expr = "static_cast<int64_t>(value.as_varint())"
|
||||
decode_varint = "static_cast<int64_t>(value)"
|
||||
encode_func = "encode_int64"
|
||||
wire_type = WireType.VARINT # Uses wire type 0
|
||||
|
||||
@@ -696,7 +679,7 @@ class UInt64Type(VarintTypeMixin, TypeInfo):
|
||||
cpp_type = "uint64_t"
|
||||
_varint_max_bits = 64
|
||||
default_value = "0"
|
||||
decode_expr = "value.as_varint()"
|
||||
decode_varint = "value"
|
||||
encode_func = "encode_uint64"
|
||||
wire_type = WireType.VARINT # Uses wire type 0
|
||||
|
||||
@@ -714,11 +697,11 @@ class UInt64Type(VarintTypeMixin, TypeInfo):
|
||||
return self._get_simple_size_calculation(name, force, "uint64")
|
||||
|
||||
@property
|
||||
def RAW_ENCODE_MAP(self) -> dict[str, tuple[str, str]]: # noqa: N802
|
||||
def RAW_ENCODE_MAP(self) -> dict[str, str]: # noqa: N802
|
||||
if self.mac_address:
|
||||
return {
|
||||
**TypeInfo.RAW_ENCODE_MAP,
|
||||
"encode_uint64": ("encode_varint_raw_48bit", "{value}"),
|
||||
"encode_uint64": "ProtoEncode::encode_varint_raw_48bit(pos, {value});",
|
||||
}
|
||||
return TypeInfo.RAW_ENCODE_MAP
|
||||
|
||||
@@ -731,7 +714,7 @@ class Int32Type(VarintTypeMixin, TypeInfo):
|
||||
cpp_type = "int32_t"
|
||||
_varint_max_bits = 64 # int32 is sign-extended to 64 bits in protobuf
|
||||
default_value = "0"
|
||||
decode_expr = "static_cast<int32_t>(value.as_varint())"
|
||||
decode_varint = "static_cast<int32_t>(value)"
|
||||
encode_func = "encode_int32"
|
||||
wire_type = WireType.VARINT # Uses wire type 0
|
||||
|
||||
@@ -751,6 +734,7 @@ class Int32Type(VarintTypeMixin, TypeInfo):
|
||||
class Fixed64Type(FixedSizeTypeMixin, TypeInfo):
|
||||
cpp_type = "uint64_t"
|
||||
default_value = "0"
|
||||
decode_64bit = "value.as_fixed64()"
|
||||
encode_func = "encode_fixed64"
|
||||
wire_type = WireType.FIXED64 # Uses wire type 1
|
||||
|
||||
@@ -776,7 +760,7 @@ class Fixed64Type(FixedSizeTypeMixin, TypeInfo):
|
||||
class Fixed32Type(FixedSizeTypeMixin, TypeInfo):
|
||||
cpp_type = "uint32_t"
|
||||
default_value = "0"
|
||||
decode_expr = "value.as_fixed32()"
|
||||
decode_32bit = "value.as_fixed32()"
|
||||
encode_func = "encode_fixed32"
|
||||
wire_type = WireType.FIXED32 # Uses wire type 5
|
||||
|
||||
@@ -785,7 +769,15 @@ class Fixed32Type(FixedSizeTypeMixin, TypeInfo):
|
||||
o += "out.append(buffer);"
|
||||
return o
|
||||
|
||||
fixed32_value_template = "{value}"
|
||||
@property
|
||||
def encode_content(self) -> str:
|
||||
tag = self.calculate_tag()
|
||||
if self.force and tag < 128:
|
||||
# Emit combined tag+value write: precomputed tag + direct memcpy
|
||||
return f"ProtoEncode::write_tag_and_fixed32(pos, {tag}, this->{self.field_name});"
|
||||
if 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:
|
||||
field_id_size = self.calculate_field_id_size()
|
||||
@@ -805,7 +797,7 @@ class BoolType(VarintTypeMixin, TypeInfo):
|
||||
_varint_max_bits = 1
|
||||
cpp_type = "bool"
|
||||
default_value = "false"
|
||||
decode_expr = "value.as_bool()"
|
||||
decode_varint = "value != 0"
|
||||
encode_func = "encode_bool"
|
||||
wire_type = WireType.VARINT # Uses wire type 0
|
||||
|
||||
@@ -825,7 +817,7 @@ class StringType(TypeInfo):
|
||||
default_value = ""
|
||||
reference_type = "std::string &"
|
||||
const_reference_type = "const std::string &"
|
||||
decode_expr = "value.as_string()"
|
||||
decode_length = "value.as_string()"
|
||||
encode_func = "encode_string"
|
||||
wire_type = WireType.LENGTH_DELIMITED # Uses wire type 2
|
||||
|
||||
@@ -859,12 +851,9 @@ class StringType(TypeInfo):
|
||||
f"this->{self.field_name}_ref_.size()",
|
||||
):
|
||||
return result
|
||||
return _encode_call(
|
||||
"encode_string",
|
||||
str(self.number),
|
||||
f"this->{self.field_name}_ref_",
|
||||
force=self.force,
|
||||
)
|
||||
if self.force:
|
||||
return f"ProtoEncode::encode_string(pos, {self.number}, this->{self.field_name}_ref_, true);"
|
||||
return f"ProtoEncode::encode_string(pos, {self.number}, this->{self.field_name}_ref_);"
|
||||
|
||||
def dump(self, name):
|
||||
# If name is 'it', this is a repeated field element - always use string
|
||||
@@ -940,9 +929,6 @@ class MessageType(TypeInfo):
|
||||
def can_use_dump_field(cls) -> bool:
|
||||
return False
|
||||
|
||||
def encode_element(self, number: int, element: str) -> str:
|
||||
return _encode_call("encode_sub_message", "buffer", str(number), element)
|
||||
|
||||
@property
|
||||
def cpp_type(self) -> str:
|
||||
return self._field.type_name[1:]
|
||||
@@ -965,9 +951,15 @@ 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}"
|
||||
)
|
||||
return f"ProtoEncode::{self.encode_func}(pos, buffer, {self.number}, this->{self.field_name});"
|
||||
|
||||
@property
|
||||
def decode_length(self) -> str:
|
||||
# Override to return None for message types because we can't use template-based
|
||||
# decoding when the specific message type isn't known at compile time.
|
||||
# Instead, we use the non-template decode_to_message() method which allows
|
||||
# runtime polymorphism through virtual function calls.
|
||||
return None
|
||||
|
||||
@property
|
||||
def public_content(self) -> list[str]:
|
||||
@@ -984,14 +976,19 @@ class MessageType(TypeInfo):
|
||||
)
|
||||
|
||||
@property
|
||||
def decode_content(self) -> str:
|
||||
body = f"value.decode_to_message(this->{self.field_name});"
|
||||
def decode_length_content(self) -> str:
|
||||
# Custom decode that doesn't use templates
|
||||
if self._track_presence:
|
||||
# decode_to_message() cannot report failure, so setting the flag
|
||||
# afterwards only documents intent; a status-returning decode could
|
||||
# gate it for real without touching callers.
|
||||
body += f"\nthis->has_{self.name} = true;"
|
||||
return self.decode_case(body)
|
||||
return (
|
||||
f"case {self.number}:\n"
|
||||
f" value.decode_to_message(this->{self.field_name});\n"
|
||||
f" this->has_{self.name} = true;\n"
|
||||
f" break;"
|
||||
)
|
||||
return f"case {self.number}: value.decode_to_message(this->{self.field_name}); break;"
|
||||
|
||||
def dump(self, name: str) -> str:
|
||||
return f"{name}.dump_to(out);"
|
||||
@@ -1030,7 +1027,7 @@ class BytesType(TypeInfo):
|
||||
reference_type = "std::string &"
|
||||
const_reference_type = "const std::string &"
|
||||
encode_func = "encode_bytes"
|
||||
decode_expr = "value.as_string()"
|
||||
decode_length = "value.as_string()"
|
||||
wire_type = WireType.LENGTH_DELIMITED # Uses wire type 2
|
||||
|
||||
@property
|
||||
@@ -1061,13 +1058,9 @@ class BytesType(TypeInfo):
|
||||
f"this->{self.field_name}_ptr_", f"this->{self.field_name}_len_"
|
||||
):
|
||||
return result
|
||||
return _encode_call(
|
||||
"encode_bytes",
|
||||
str(self.number),
|
||||
f"this->{self.field_name}_ptr_",
|
||||
f"this->{self.field_name}_len_",
|
||||
force=self.force,
|
||||
)
|
||||
if self.force:
|
||||
return f"ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}_ptr_, this->{self.field_name}_len_, true);"
|
||||
return f"ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}_ptr_, this->{self.field_name}_len_);"
|
||||
|
||||
def dump(self, name: str) -> str:
|
||||
ptr_dump = f"format_hex_pretty(this->{self.field_name}_ptr_, this->{self.field_name}_len_)"
|
||||
@@ -1140,6 +1133,11 @@ class PointerToBufferTypeBase(TypeInfo):
|
||||
super().__init__(field)
|
||||
self.array_size = 0
|
||||
|
||||
@property
|
||||
def decode_length(self) -> str | None:
|
||||
# This is handled in decode_length_content
|
||||
return None
|
||||
|
||||
@property
|
||||
def wire_type(self) -> WireType:
|
||||
"""Get the wire type for this field."""
|
||||
@@ -1172,20 +1170,17 @@ class PointerToBytesBufferType(PointerToBufferTypeBase):
|
||||
f"this->{self.field_name}", f"this->{self.field_name}_len"
|
||||
):
|
||||
return result
|
||||
return _encode_call(
|
||||
"encode_bytes",
|
||||
str(self.number),
|
||||
f"this->{self.field_name}",
|
||||
f"this->{self.field_name}_len",
|
||||
force=self.force,
|
||||
)
|
||||
if self.force:
|
||||
return f"ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}, this->{self.field_name}_len, true);"
|
||||
return f"ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}, this->{self.field_name}_len);"
|
||||
|
||||
@property
|
||||
def decode_content(self) -> str:
|
||||
return self.decode_case(
|
||||
f"this->{self.field_name} = value.data();\n"
|
||||
f"this->{self.field_name}_len = value.size();",
|
||||
)
|
||||
def decode_length_content(self) -> str | None:
|
||||
return f"""case {self.number}: {{
|
||||
this->{self.field_name} = value.data();
|
||||
this->{self.field_name}_len = value.size();
|
||||
break;
|
||||
}}"""
|
||||
|
||||
def dump(self, name: str) -> str:
|
||||
return (
|
||||
@@ -1229,26 +1224,24 @@ class PointerToStringBufferType(PointerToBufferTypeBase):
|
||||
if max_len is not None and max_len < 128 and self.force:
|
||||
tag = self.calculate_tag()
|
||||
if tag < 128:
|
||||
return _encode_call(
|
||||
"encode_short_string_force", str(tag), f"this->{self.field_name}"
|
||||
)
|
||||
return f"ProtoEncode::encode_short_string_force(pos, {tag}, this->{self.field_name});"
|
||||
if result := self._encode_bytes_with_precomputed_tag(
|
||||
f"this->{self.field_name}.c_str()",
|
||||
f"this->{self.field_name}.size()",
|
||||
):
|
||||
return result
|
||||
return _encode_call(
|
||||
"encode_string",
|
||||
str(self.number),
|
||||
f"this->{self.field_name}",
|
||||
force=self.force,
|
||||
if self.force:
|
||||
return f"ProtoEncode::encode_string(pos, {self.number}, this->{self.field_name}, true);"
|
||||
return (
|
||||
f"ProtoEncode::encode_string(pos, {self.number}, this->{self.field_name});"
|
||||
)
|
||||
|
||||
@property
|
||||
def decode_content(self) -> str:
|
||||
return self.decode_case(
|
||||
f"this->{self.field_name} = StringRef(value.data(), value.size());",
|
||||
)
|
||||
def decode_length_content(self) -> str | None:
|
||||
return f"""case {self.number}: {{
|
||||
this->{self.field_name} = StringRef(reinterpret_cast<const char *>(value.data()), value.size());
|
||||
break;
|
||||
}}"""
|
||||
|
||||
def dump(self, name: str) -> str:
|
||||
# Not used since we use dump_field, but required by abstract base class
|
||||
@@ -1317,13 +1310,14 @@ class PackedBufferTypeInfo(TypeInfo):
|
||||
]
|
||||
|
||||
@property
|
||||
def decode_content(self) -> str:
|
||||
def decode_length_content(self) -> str:
|
||||
"""Store pointer to buffer and calculate count of packed varints."""
|
||||
return self.decode_case(
|
||||
f"this->{self.field_name}_data_ = value.data();\n"
|
||||
f"this->{self.field_name}_length_ = value.size();\n"
|
||||
f"this->{self.field_name}_count_ = count_packed_varints(value.data(), value.size());",
|
||||
)
|
||||
return f"""case {self.number}: {{
|
||||
this->{self.field_name}_data_ = value.data();
|
||||
this->{self.field_name}_length_ = value.size();
|
||||
this->{self.field_name}_count_ = count_packed_varints(value.data(), value.size());
|
||||
break;
|
||||
}}"""
|
||||
|
||||
@property
|
||||
def encode_content(self) -> str:
|
||||
@@ -1408,11 +1402,17 @@ class FixedArrayBytesType(TypeInfo):
|
||||
]
|
||||
|
||||
@property
|
||||
def decode_content(self) -> str:
|
||||
return self.decode_case(
|
||||
f"this->{self.field_name}_len = std::min<size_t>(value.size(), {self.array_size});\n"
|
||||
f"memcpy(this->{self.field_name}, value.data(), this->{self.field_name}_len);",
|
||||
)
|
||||
def decode_length_content(self) -> str:
|
||||
o = f"case {self.number}: {{\n"
|
||||
o += " const std::string &data_str = value.as_string();\n"
|
||||
o += f" this->{self.field_name}_len = data_str.size();\n"
|
||||
o += f" if (this->{self.field_name}_len > {self.array_size}) {{\n"
|
||||
o += f" this->{self.field_name}_len = {self.array_size};\n"
|
||||
o += " }\n"
|
||||
o += f" memcpy(this->{self.field_name}, data_str.data(), this->{self.field_name}_len);\n"
|
||||
o += " break;\n"
|
||||
o += "}"
|
||||
return o
|
||||
|
||||
@property
|
||||
def encode_content(self) -> str:
|
||||
@@ -1421,13 +1421,9 @@ class FixedArrayBytesType(TypeInfo):
|
||||
f"this->{self.field_name}", f"this->{self.field_name}_len", max_len=max_len
|
||||
):
|
||||
return result
|
||||
return _encode_call(
|
||||
"encode_bytes",
|
||||
str(self.number),
|
||||
f"this->{self.field_name}",
|
||||
f"this->{self.field_name}_len",
|
||||
force=self.force,
|
||||
)
|
||||
if self.force:
|
||||
return f"ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}, this->{self.field_name}_len, true);"
|
||||
return f"ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}, this->{self.field_name}_len);"
|
||||
|
||||
def dump(self, name: str) -> str:
|
||||
return f"out.append(format_hex_pretty({name}, {name}_len));"
|
||||
@@ -1475,7 +1471,7 @@ class UInt32Type(VarintTypeMixin, TypeInfo):
|
||||
cpp_type = "uint32_t"
|
||||
_varint_max_bits = 32
|
||||
default_value = "0"
|
||||
decode_expr = "value.as_varint()"
|
||||
decode_varint = "value"
|
||||
encode_func = "encode_uint32"
|
||||
wire_type = WireType.VARINT # Uses wire type 0
|
||||
|
||||
@@ -1498,21 +1494,13 @@ class UInt32Type(VarintTypeMixin, TypeInfo):
|
||||
class EnumType(VarintTypeMixin, TypeInfo):
|
||||
_varint_max_bits = 32
|
||||
|
||||
def encode_element(self, number: int, element: str) -> str:
|
||||
return _encode_call(
|
||||
self.encode_func,
|
||||
str(number),
|
||||
f"static_cast<uint32_t>({element})",
|
||||
force=True,
|
||||
)
|
||||
|
||||
@property
|
||||
def cpp_type(self) -> str:
|
||||
return f"enums::{self._field.type_name[1:]}"
|
||||
|
||||
@property
|
||||
def decode_expr(self) -> str:
|
||||
return f"static_cast<{self.cpp_type}>(value.as_varint())"
|
||||
def decode_varint(self) -> str:
|
||||
return f"static_cast<{self.cpp_type}>(value)"
|
||||
|
||||
default_value = ""
|
||||
wire_type = WireType.VARINT # Uses wire type 0
|
||||
@@ -1532,9 +1520,9 @@ class EnumType(VarintTypeMixin, TypeInfo):
|
||||
@property
|
||||
def encode_content(self) -> str:
|
||||
value_expr = f"static_cast<uint32_t>(this->{self.field_name})"
|
||||
return _encode_call(
|
||||
self.encode_func, str(self.number), value_expr, force=self.force
|
||||
)
|
||||
if self.force:
|
||||
return f"ProtoEncode::{self.encode_func}(pos, {self.number}, {value_expr}, true);"
|
||||
return f"ProtoEncode::{self.encode_func}(pos, {self.number}, {value_expr});"
|
||||
|
||||
def dump(self, name: str) -> str:
|
||||
return f"out.append_p(proto_enum_to_string<{self.cpp_type}>({name}));"
|
||||
@@ -1559,7 +1547,7 @@ class EnumType(VarintTypeMixin, TypeInfo):
|
||||
class SFixed32Type(FixedSizeTypeMixin, TypeInfo):
|
||||
cpp_type = "int32_t"
|
||||
default_value = "0"
|
||||
decode_expr = "value.as_sfixed32()"
|
||||
decode_32bit = "value.as_sfixed32()"
|
||||
encode_func = "encode_sfixed32"
|
||||
wire_type = WireType.FIXED32 # Uses wire type 5
|
||||
|
||||
@@ -1585,6 +1573,7 @@ class SFixed32Type(FixedSizeTypeMixin, TypeInfo):
|
||||
class SFixed64Type(FixedSizeTypeMixin, TypeInfo):
|
||||
cpp_type = "int64_t"
|
||||
default_value = "0"
|
||||
decode_64bit = "value.as_sfixed64()"
|
||||
encode_func = "encode_sfixed64"
|
||||
wire_type = WireType.FIXED64 # Uses wire type 1
|
||||
|
||||
@@ -1611,7 +1600,7 @@ class SInt32Type(VarintTypeMixin, TypeInfo):
|
||||
cpp_type = "int32_t"
|
||||
_varint_max_bits = 32 # zigzag encoding keeps it 32-bit
|
||||
default_value = "0"
|
||||
decode_expr = "decode_zigzag32(static_cast<uint32_t>(value.as_varint()))"
|
||||
decode_varint = "decode_zigzag32(static_cast<uint32_t>(value))"
|
||||
encode_func = "encode_sint32"
|
||||
wire_type = WireType.VARINT # Uses wire type 0
|
||||
|
||||
@@ -1632,7 +1621,7 @@ class SInt64Type(VarintTypeMixin, TypeInfo):
|
||||
cpp_type = "int64_t"
|
||||
_varint_max_bits = 64
|
||||
default_value = "0"
|
||||
decode_expr = "decode_zigzag64(value.as_varint())"
|
||||
decode_varint = "decode_zigzag64(value)"
|
||||
encode_func = "encode_sint64"
|
||||
wire_type = WireType.VARINT # Uses wire type 0
|
||||
|
||||
@@ -1712,9 +1701,9 @@ def _generate_inline_encode_block(
|
||||
|
||||
lines = []
|
||||
lines.append(f"auto &sub_msg = {element};")
|
||||
lines.append(_encode_call("write_raw_byte", str(tag)))
|
||||
lines.append(f"ProtoEncode::write_raw_byte(pos, {tag});")
|
||||
lines.append("uint8_t *len_pos = pos;")
|
||||
lines.append(_encode_call("reserve_byte"))
|
||||
lines.append("ProtoEncode::reserve_byte(pos);")
|
||||
|
||||
# Generate inline field encoding for each sub-message field
|
||||
for field in sub_desc.field:
|
||||
@@ -1785,11 +1774,18 @@ class FixedArrayRepeatedType(TypeInfo):
|
||||
|
||||
def _encode_element(self, element: str) -> str:
|
||||
"""Helper to generate encode statement for a single element."""
|
||||
if isinstance(self._ti, MessageType) and _is_inline_encode(self._ti.cpp_type):
|
||||
return _generate_inline_encode_block(
|
||||
self.number, self._ti.cpp_type, element
|
||||
)
|
||||
return self._ti.encode_element(self.number, element)
|
||||
if isinstance(self._ti, EnumType):
|
||||
return f"ProtoEncode::{self._ti.encode_func}(pos, {self.number}, static_cast<uint32_t>({element}), true);"
|
||||
# Repeated message elements use encode_sub_message (force=true is default)
|
||||
if isinstance(self._ti, MessageType):
|
||||
if _is_inline_encode(self._ti.cpp_type):
|
||||
return _generate_inline_encode_block(
|
||||
self.number, self._ti.cpp_type, element
|
||||
)
|
||||
return f"ProtoEncode::encode_sub_message(pos, buffer, {self.number}, {element});"
|
||||
return (
|
||||
f"ProtoEncode::{self._ti.encode_func}(pos, {self.number}, {element}, true);"
|
||||
)
|
||||
|
||||
@property
|
||||
def cpp_type(self) -> str:
|
||||
@@ -2083,23 +2079,55 @@ class RepeatedTypeInfo(TypeInfo):
|
||||
return self._ti.wire_type
|
||||
|
||||
@property
|
||||
def decode_expr(self) -> str | None:
|
||||
return self._ti.decode_expr
|
||||
|
||||
def _decode_store(self, expr: str) -> str:
|
||||
return f"this->{self.field_name}.push_back({expr});"
|
||||
|
||||
@property
|
||||
def decode_content(self) -> str | None:
|
||||
def decode_varint_content(self) -> str:
|
||||
# Pointer fields don't support decoding
|
||||
if self._use_pointer:
|
||||
return None
|
||||
if isinstance(self._ti, MessageType):
|
||||
return self.decode_case(
|
||||
f"this->{self.field_name}.emplace_back();\n"
|
||||
f"value.decode_to_message(this->{self.field_name}.back());"
|
||||
)
|
||||
return super().decode_content
|
||||
content = self._ti.decode_varint
|
||||
if content is None:
|
||||
return None
|
||||
return (
|
||||
f"case {self.number}: this->{self.field_name}.push_back({content}); break;"
|
||||
)
|
||||
|
||||
@property
|
||||
def decode_length_content(self) -> str:
|
||||
# Pointer fields don't support decoding
|
||||
if self._use_pointer:
|
||||
return None
|
||||
content = self._ti.decode_length
|
||||
if content is None and isinstance(self._ti, MessageType):
|
||||
# Special handling for non-template message decoding
|
||||
return f"case {self.number}: this->{self.field_name}.emplace_back(); value.decode_to_message(this->{self.field_name}.back()); break;"
|
||||
if content is None:
|
||||
return None
|
||||
return (
|
||||
f"case {self.number}: this->{self.field_name}.push_back({content}); break;"
|
||||
)
|
||||
|
||||
@property
|
||||
def decode_32bit_content(self) -> str:
|
||||
# Pointer fields don't support decoding
|
||||
if self._use_pointer:
|
||||
return None
|
||||
content = self._ti.decode_32bit
|
||||
if content is None:
|
||||
return None
|
||||
return (
|
||||
f"case {self.number}: this->{self.field_name}.push_back({content}); break;"
|
||||
)
|
||||
|
||||
@property
|
||||
def decode_64bit_content(self) -> str:
|
||||
# Pointer fields don't support decoding
|
||||
if self._use_pointer:
|
||||
return None
|
||||
content = self._ti.decode_64bit
|
||||
if content is None:
|
||||
return None
|
||||
return (
|
||||
f"case {self.number}: this->{self.field_name}.push_back({content}); break;"
|
||||
)
|
||||
|
||||
@property
|
||||
def _ti_is_bool(self) -> bool:
|
||||
@@ -2107,7 +2135,15 @@ class RepeatedTypeInfo(TypeInfo):
|
||||
return isinstance(self._ti, BoolType)
|
||||
|
||||
def _encode_element_call(self, element: str) -> str:
|
||||
return self._ti.encode_element(self.number, element)
|
||||
"""Helper to generate encode call for a single element."""
|
||||
if isinstance(self._ti, EnumType):
|
||||
return f"ProtoEncode::{self._ti.encode_func}(pos, {self.number}, static_cast<uint32_t>({element}), true);"
|
||||
# Repeated message elements use encode_sub_message (force=true is default)
|
||||
if isinstance(self._ti, MessageType):
|
||||
return f"ProtoEncode::encode_sub_message(pos, buffer, {self.number}, {element});"
|
||||
return (
|
||||
f"ProtoEncode::{self._ti.encode_func}(pos, {self.number}, {element}, true);"
|
||||
)
|
||||
|
||||
@property
|
||||
def encode_content(self) -> str:
|
||||
@@ -2116,7 +2152,7 @@ class RepeatedTypeInfo(TypeInfo):
|
||||
# Special handling for const char* elements (when container_no_template contains "const char")
|
||||
if "const char" in self._container_no_template:
|
||||
o = f"for (const char *it : *this->{self.field_name}) {{\n"
|
||||
o += f" {_encode_call(self._ti.encode_func, str(self.number), 'it', 'strlen(it)', force=True)}\n"
|
||||
o += f" ProtoEncode::{self._ti.encode_func}(pos, {self.number}, it, strlen(it), true);\n"
|
||||
else:
|
||||
o = f"for (const auto &it : *this->{self.field_name}) {{\n"
|
||||
o += f" {self._encode_element_call('it')}\n"
|
||||
@@ -2502,7 +2538,10 @@ def build_message_type(
|
||||
) -> tuple[str, str, str]:
|
||||
public_content: list[str] = []
|
||||
protected_content: list[str] = []
|
||||
decode: list[str] = []
|
||||
decode_varint: list[str] = []
|
||||
decode_length: list[str] = []
|
||||
decode_32bit: list[str] = []
|
||||
decode_64bit: list[str] = []
|
||||
encode: list[str] = []
|
||||
dump: list[str] = []
|
||||
size_calc: list[str] = []
|
||||
@@ -2631,8 +2670,22 @@ def build_message_type(
|
||||
if field.options.HasExtension(pb.field_ifdef):
|
||||
field_ifdef = field.options.Extensions[pb.field_ifdef]
|
||||
|
||||
if case := ti.decode_content:
|
||||
decode.extend(wrap_with_ifdef(case, field_ifdef))
|
||||
if ti.decode_varint_content:
|
||||
decode_varint.extend(
|
||||
wrap_with_ifdef(ti.decode_varint_content, field_ifdef)
|
||||
)
|
||||
if ti.decode_length_content:
|
||||
decode_length.extend(
|
||||
wrap_with_ifdef(ti.decode_length_content, field_ifdef)
|
||||
)
|
||||
if ti.decode_32bit_content:
|
||||
decode_32bit.extend(
|
||||
wrap_with_ifdef(ti.decode_32bit_content, field_ifdef)
|
||||
)
|
||||
if ti.decode_64bit_content:
|
||||
decode_64bit.extend(
|
||||
wrap_with_ifdef(ti.decode_64bit_content, field_ifdef)
|
||||
)
|
||||
if ti.dump_content:
|
||||
# Check for field_ifdef option for dump as well
|
||||
field_ifdef = None
|
||||
@@ -2642,15 +2695,49 @@ def build_message_type(
|
||||
dump.extend(wrap_with_ifdef(ti.dump_content, field_ifdef))
|
||||
|
||||
cpp = ""
|
||||
if decode:
|
||||
o = f"void {desc.name}::decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) {{\n"
|
||||
o += " const ProtoFieldValue value(data, scalar);\n"
|
||||
o += " switch (tag) {\n"
|
||||
o += indent("\n".join(decode), " ") + "\n"
|
||||
if decode_varint:
|
||||
o = f"bool {desc.name}::decode_varint(uint32_t field_id, proto_varint_value_t value) {{\n"
|
||||
o += " switch (field_id) {\n"
|
||||
o += indent("\n".join(decode_varint), " ") + "\n"
|
||||
o += " default: return false;\n"
|
||||
o += " }\n"
|
||||
o += " return true;\n"
|
||||
o += "}\n"
|
||||
cpp += o
|
||||
prot = "void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;"
|
||||
prot = "bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;"
|
||||
protected_content.insert(0, prot)
|
||||
if decode_length:
|
||||
o = f"bool {desc.name}::decode_length(uint32_t field_id, ProtoLengthDelimited value) {{\n"
|
||||
o += " switch (field_id) {\n"
|
||||
o += indent("\n".join(decode_length), " ") + "\n"
|
||||
o += " default: return false;\n"
|
||||
o += " }\n"
|
||||
o += " return true;\n"
|
||||
o += "}\n"
|
||||
cpp += o
|
||||
prot = "bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;"
|
||||
protected_content.insert(0, prot)
|
||||
if decode_32bit:
|
||||
o = f"bool {desc.name}::decode_32bit(uint32_t field_id, Proto32Bit value) {{\n"
|
||||
o += " switch (field_id) {\n"
|
||||
o += indent("\n".join(decode_32bit), " ") + "\n"
|
||||
o += " default: return false;\n"
|
||||
o += " }\n"
|
||||
o += " return true;\n"
|
||||
o += "}\n"
|
||||
cpp += o
|
||||
prot = "bool decode_32bit(uint32_t field_id, Proto32Bit value) override;"
|
||||
protected_content.insert(0, prot)
|
||||
if decode_64bit:
|
||||
o = f"bool {desc.name}::decode_64bit(uint32_t field_id, Proto64Bit value) {{\n"
|
||||
o += " switch (field_id) {\n"
|
||||
o += indent("\n".join(decode_64bit), " ") + "\n"
|
||||
o += " default: return false;\n"
|
||||
o += " }\n"
|
||||
o += " return true;\n"
|
||||
o += "}\n"
|
||||
cpp += o
|
||||
prot = "bool decode_64bit(uint32_t field_id, Proto64Bit value) override;"
|
||||
protected_content.insert(0, prot)
|
||||
|
||||
# Generate custom decode() override for messages with FixedVector fields
|
||||
@@ -2697,36 +2784,28 @@ 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" const auto &msg = *static_cast<const {desc.name} *>(self);\n"
|
||||
o = f"{speed_attr}uint8_t *{desc.name}::encode(ProtoWriteBuffer &buffer PROTO_ENCODE_DEBUG_PARAM) const {{\n"
|
||||
o += " uint8_t *__restrict__ pos = buffer.get_pos();\n"
|
||||
o += indent("\n".join(encode_debug)).replace("this->", "msg.") + "\n"
|
||||
o += indent("\n".join(encode_debug)) + "\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);"
|
||||
)
|
||||
public_content.append(
|
||||
"uint8_t *encode(ProtoWriteBuffer &buffer PROTO_ENCODE_DEBUG_PARAM) const {\n"
|
||||
" return encode_msg(this, buffer PROTO_ENCODE_DEBUG_ARG);\n"
|
||||
"}"
|
||||
prot = (
|
||||
"uint8_t *encode(ProtoWriteBuffer &buffer PROTO_ENCODE_DEBUG_PARAM) const;"
|
||||
)
|
||||
public_content.append(prot)
|
||||
# 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
|
||||
if needs_encode and size_calc and not is_inline_only:
|
||||
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 = f"{speed_attr}uint32_t {desc.name}::calculate_size() const {{\n"
|
||||
o += " uint32_t size = 0;\n"
|
||||
o += indent("\n".join(size_calc)).replace("this->", "msg.") + "\n"
|
||||
o += indent("\n".join(size_calc)) + "\n"
|
||||
o += " return size;\n"
|
||||
o += "}\n"
|
||||
cpp += o
|
||||
public_content.append("static uint32_t calc_size_msg(const void *self);")
|
||||
public_content.append(
|
||||
"uint32_t calculate_size() const { return calc_size_msg(this); }"
|
||||
)
|
||||
prot = "uint32_t calculate_size() const;"
|
||||
public_content.append(prot)
|
||||
# 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
|
||||
|
||||
@@ -249,7 +249,7 @@ static APIBuffer build_infrared_rf_transmit_wire() {
|
||||
std::memcpy(bytes + len, packed, packed_len);
|
||||
len += packed_len;
|
||||
// field 6: modulation = 1 (non-zero so it's actually emitted and exercises
|
||||
// decode_field for this field, matching the documented layout above).
|
||||
// decode_varint for this field, matching the documented layout above).
|
||||
put_byte(0x30);
|
||||
put_varint(1);
|
||||
|
||||
|
||||
@@ -0,0 +1,9 @@
|
||||
esphome:
|
||||
name: test
|
||||
|
||||
esp32:
|
||||
board: esp32dev
|
||||
|
||||
remote_receiver:
|
||||
- id: rcvr
|
||||
pin: GPIO4
|
||||
@@ -0,0 +1,24 @@
|
||||
esphome:
|
||||
name: test
|
||||
|
||||
esp32:
|
||||
board: esp32dev
|
||||
|
||||
logger:
|
||||
|
||||
remote_receiver:
|
||||
- id: rcvr
|
||||
pin: GPIO4
|
||||
dump:
|
||||
- nec
|
||||
- rc_switch
|
||||
on_nec:
|
||||
then:
|
||||
- logger.log: nec
|
||||
|
||||
binary_sensor:
|
||||
- platform: remote_receiver
|
||||
name: Remote Input
|
||||
nec:
|
||||
address: 0x1234
|
||||
command: 0x5678
|
||||
@@ -0,0 +1,20 @@
|
||||
esphome:
|
||||
name: test
|
||||
|
||||
esp32:
|
||||
board: esp32dev
|
||||
|
||||
remote_receiver:
|
||||
- id: rcvr
|
||||
pin: GPIO4
|
||||
|
||||
infrared:
|
||||
- platform: ir_rf_proxy
|
||||
name: IR Receiver
|
||||
remote_receiver_id: rcvr
|
||||
|
||||
radio_frequency:
|
||||
- platform: ir_rf_proxy
|
||||
name: RF Receiver
|
||||
frequency: 433.92MHz
|
||||
remote_receiver_id: rcvr
|
||||
@@ -0,0 +1,82 @@
|
||||
"""Listener and dumper StaticVector sizes come from codegen slot counts."""
|
||||
|
||||
from collections.abc import Callable
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from esphome.components import remote_base
|
||||
import esphome.config_validation as cv
|
||||
|
||||
from ..helpers import get_define_value
|
||||
|
||||
|
||||
def test_dumper_and_listener_counts(
|
||||
generate_main: Callable[[str | Path], str],
|
||||
component_config_path: Callable[[str], Path],
|
||||
) -> None:
|
||||
generate_main(component_config_path("receiver_with_dumpers.yaml"))
|
||||
# nec and rc_switch dumpers
|
||||
assert get_define_value("REMOTE_BASE_DUMPER_COUNT") == "2"
|
||||
# on_nec trigger plus the remote_receiver binary sensor
|
||||
assert get_define_value("REMOTE_BASE_LISTENER_COUNT") == "2"
|
||||
|
||||
|
||||
def test_bare_receiver_emits_no_counts(
|
||||
generate_main: Callable[[str | Path], str],
|
||||
component_config_path: Callable[[str], Path],
|
||||
) -> None:
|
||||
generate_main(component_config_path("receiver_bare.yaml"))
|
||||
assert get_define_value("REMOTE_BASE_DUMPER_COUNT") is None
|
||||
assert get_define_value("REMOTE_BASE_LISTENER_COUNT") is None
|
||||
|
||||
|
||||
def test_proxy_receivers_count_as_listeners(
|
||||
generate_main: Callable[[str | Path], str],
|
||||
component_config_path: Callable[[str], Path],
|
||||
) -> None:
|
||||
generate_main(component_config_path("receiver_with_proxies.yaml"))
|
||||
# infrared and radio_frequency ir_rf_proxy platforms each listen
|
||||
assert get_define_value("REMOTE_BASE_LISTENER_COUNT") == "2"
|
||||
assert get_define_value("REMOTE_BASE_DUMPER_COUNT") is None
|
||||
|
||||
|
||||
def test_only_used_protocol_sources_are_compiled(
|
||||
generate_main: Callable[[str | Path], str],
|
||||
component_config_path: Callable[[str], Path],
|
||||
) -> None:
|
||||
generate_main(component_config_path("receiver_with_dumpers.yaml"))
|
||||
excluded = set(remote_base.FILTER_SOURCE_FILES())
|
||||
assert "nec_protocol.cpp" not in excluded
|
||||
assert "rc_switch_protocol.cpp" not in excluded
|
||||
assert "sony_protocol.cpp" in excluded
|
||||
assert "remote_base.cpp" not in excluded
|
||||
|
||||
|
||||
def test_every_registry_name_maps_to_a_protocol_source() -> None:
|
||||
sources = {
|
||||
path.name for path in Path(remote_base.__file__).parent.glob("*_protocol.cpp")
|
||||
}
|
||||
names = (
|
||||
set(remote_base.BINARY_SENSOR_REGISTRY)
|
||||
| set(remote_base.DUMPER_REGISTRY)
|
||||
| {key.removeprefix("on_") for key in remote_base.TRIGGER_REGISTRY}
|
||||
)
|
||||
for name in names:
|
||||
stem = (
|
||||
remote_base.protocol_define(name)
|
||||
.removeprefix("USE_REMOTE_PROTOCOL_")
|
||||
.lower()
|
||||
)
|
||||
assert f"{stem}_protocol.cpp" in sources, name
|
||||
|
||||
|
||||
def test_dump_list_is_deduplicated_across_forms() -> None:
|
||||
dumpers = remote_base.validate_dumpers(["raw", {"raw": None}, "nec", "nec"])
|
||||
assert [name for name, _ in dumpers] == ["raw", "nec"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("bad", [["nec", None], [5]])
|
||||
def test_dump_list_rejects_invalid_entries_with_a_validation_error(bad: list) -> None:
|
||||
with pytest.raises(cv.Invalid):
|
||||
remote_base.validate_dumpers(bad)
|
||||
@@ -59,7 +59,7 @@ static void verify_mac(uint64_t mac, size_t expected_bytes) {
|
||||
#ifdef ESPHOME_DEBUG_API
|
||||
uint8_t *proto_debug_end_ = api_buf.data() + api_buf.size();
|
||||
#endif
|
||||
pos = ProtoEncode::encode_varint_raw_48bit(pos PROTO_ENCODE_DEBUG_ARG, mac);
|
||||
ProtoEncode::encode_varint_raw_48bit(pos PROTO_ENCODE_DEBUG_ARG, mac);
|
||||
size_t new_len = pos - api_buf.data();
|
||||
|
||||
EXPECT_EQ(new_len, expected_bytes) << "mac=0x" << std::hex << mac << std::dec;
|
||||
|
||||
@@ -0,0 +1,6 @@
|
||||
# A receiver with no dumpers and no listeners compiles both lists out.
|
||||
# Only built while remote_receiver is tested in isolation: the counts are global defines,
|
||||
# so this variant cannot be merged with configs that register any.
|
||||
remote_receiver:
|
||||
- id: rcvr_bare
|
||||
pin: ${pin}
|
||||
@@ -0,0 +1,5 @@
|
||||
substitutions:
|
||||
pin: GPIO2
|
||||
|
||||
packages:
|
||||
bare: !include bare-common.yaml
|
||||
@@ -1,43 +0,0 @@
|
||||
esphome:
|
||||
name: api-decode-wire-types-test
|
||||
host:
|
||||
api:
|
||||
logger:
|
||||
level: DEBUG
|
||||
|
||||
switch:
|
||||
- platform: template
|
||||
name: "Wire Switch"
|
||||
optimistic: true
|
||||
|
||||
output:
|
||||
- platform: template
|
||||
id: wire_dim
|
||||
type: float
|
||||
write_action:
|
||||
- lambda: ""
|
||||
|
||||
light:
|
||||
- platform: monochromatic
|
||||
name: "Wire Light"
|
||||
output: wire_dim
|
||||
default_transition_length: 0s
|
||||
effects:
|
||||
- pulse:
|
||||
name: Pulse
|
||||
|
||||
text:
|
||||
- platform: template
|
||||
name: "Wire Text"
|
||||
optimistic: true
|
||||
mode: text
|
||||
min_length: 0
|
||||
max_length: 255
|
||||
|
||||
number:
|
||||
- platform: template
|
||||
name: "Wire Number"
|
||||
optimistic: true
|
||||
min_value: -1000
|
||||
max_value: 1000
|
||||
step: 0.5
|
||||
@@ -1,11 +0,0 @@
|
||||
esphome:
|
||||
name: api-empty-message-test
|
||||
host:
|
||||
api:
|
||||
logger:
|
||||
level: DEBUG
|
||||
|
||||
switch:
|
||||
- platform: template
|
||||
name: "Empty Message Switch"
|
||||
optimistic: true
|
||||
@@ -1,58 +0,0 @@
|
||||
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');
|
||||
@@ -125,12 +125,11 @@ class RawApiClient:
|
||||
await self.read_until_frame(MESSAGE_TYPE_OF[api_pb2.HelloResponse])
|
||||
|
||||
async def send_message(self, msg: message.Message) -> None:
|
||||
await self.send_raw(MESSAGE_TYPE_OF[type(msg)], msg.SerializeToString())
|
||||
|
||||
async def send_raw(self, msg_type: int, payload: bytes) -> None:
|
||||
"""Send a frame with a hand built payload, for shapes protobuf will not serialize."""
|
||||
loop = asyncio.get_running_loop()
|
||||
await loop.sock_sendall(self._sock, encode_frame(msg_type, payload))
|
||||
await loop.sock_sendall(
|
||||
self._sock,
|
||||
encode_frame(MESSAGE_TYPE_OF[type(msg)], msg.SerializeToString()),
|
||||
)
|
||||
|
||||
async def read_until_frame(self, msg_type: int, timeout: float = 10.0) -> None:
|
||||
"""Read until at least one frame of msg_type has been received."""
|
||||
|
||||
@@ -57,46 +57,6 @@ async def wait_for_state(
|
||||
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](
|
||||
entities: list[EntityInfo],
|
||||
object_id_substring: str,
|
||||
|
||||
@@ -1,142 +0,0 @@
|
||||
"""decode_field() must take fields that match their declared wire type, drop the ones that do
|
||||
not, skip unknown fields, and handle two byte tags, varints and length prefixes."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
import struct
|
||||
|
||||
from aioesphomeapi import (
|
||||
EntityState,
|
||||
LightState,
|
||||
NumberState,
|
||||
SwitchState,
|
||||
TextState,
|
||||
api_pb2,
|
||||
)
|
||||
import pytest
|
||||
|
||||
from .raw_api_client import MESSAGE_TYPE_OF, RawApiClient, encode_varint
|
||||
from .state_utils import InitialStateHelper, StateWaiter, require_entity
|
||||
from .types import APIClientConnectedFactory, RunCompiledFunction
|
||||
|
||||
SWITCH_COMMAND = MESSAGE_TYPE_OF[api_pb2.SwitchCommandRequest]
|
||||
WIRE_VARINT, WIRE_LENGTH, WIRE_FIXED32 = 0, 2, 5
|
||||
|
||||
|
||||
def tag(field: int, wire_type: int) -> bytes:
|
||||
return encode_varint((field << 3) | wire_type)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_api_decode_wire_types(
|
||||
yaml_config: str,
|
||||
run_compiled: RunCompiledFunction,
|
||||
api_client_connected: APIClientConnectedFactory,
|
||||
unused_tcp_port: int,
|
||||
) -> None:
|
||||
async with (
|
||||
run_compiled(yaml_config),
|
||||
api_client_connected() as client,
|
||||
RawApiClient(unused_tcp_port) as raw,
|
||||
):
|
||||
entities, _ = await client.list_entities_services()
|
||||
switch = require_entity(entities, "wire_switch")
|
||||
light = require_entity(entities, "wire_light")
|
||||
text = require_entity(entities, "wire_text")
|
||||
number = require_entity(entities, "wire_number")
|
||||
key = tag(1, WIRE_FIXED32) + struct.pack("<I", switch.key)
|
||||
on, off = tag(2, WIRE_VARINT) + b"\x01", tag(2, WIRE_VARINT) + b"\x00"
|
||||
|
||||
switch_states: list[bool] = []
|
||||
waiter = StateWaiter()
|
||||
|
||||
def on_state(state: EntityState) -> None:
|
||||
if isinstance(state, SwitchState) and state.key == switch.key:
|
||||
switch_states.append(state.state)
|
||||
waiter.on_state(state)
|
||||
|
||||
def switch_is(value: bool) -> Callable[[EntityState], bool]:
|
||||
return lambda s: (
|
||||
isinstance(s, SwitchState) and s.key == switch.key and s.state is value
|
||||
)
|
||||
|
||||
def number_is(value: float) -> Callable[[EntityState], bool]:
|
||||
return lambda s: (
|
||||
isinstance(s, NumberState) and s.key == number.key and s.state == value
|
||||
)
|
||||
|
||||
initial = InitialStateHelper(entities)
|
||||
client.subscribe_states(initial.on_state_wrapper(on_state))
|
||||
await initial.wait_for_initial_states()
|
||||
await raw.connect()
|
||||
|
||||
# A well formed command: fixed32 key, varint state
|
||||
await raw.send_raw(SWITCH_COMMAND, key + on)
|
||||
await waiter.expect(switch_is(True))
|
||||
await raw.send_raw(SWITCH_COMMAND, key + off)
|
||||
await waiter.expect(switch_is(False))
|
||||
|
||||
# The same field with the wrong wire type is dropped, and a varint key never matches an
|
||||
# entity; each of these would turn the switch on if the payload were read as a varint
|
||||
seen = len(switch_states)
|
||||
await raw.send_raw(SWITCH_COMMAND, key + tag(2, WIRE_LENGTH) + b"\x01\x01")
|
||||
await raw.send_raw(
|
||||
SWITCH_COMMAND, key + tag(2, WIRE_FIXED32) + b"\x01\x00\x00\x00"
|
||||
)
|
||||
await raw.send_raw(
|
||||
SWITCH_COMMAND, tag(1, WIRE_VARINT) + encode_varint(switch.key) + on
|
||||
)
|
||||
# Ordered on the raw socket itself: this frame cannot be parsed before the bad ones, so
|
||||
# the only switch state since the marker must be the one it produces
|
||||
await raw.send_raw(SWITCH_COMMAND, key + on)
|
||||
await waiter.expect(switch_is(True), label="switch on after wrong wire types")
|
||||
assert switch_states[seen:] == [True]
|
||||
await raw.send_raw(SWITCH_COMMAND, key + off)
|
||||
await waiter.expect(switch_is(False))
|
||||
|
||||
# Truncated bodies stop the decode loop without taking the connection down: a tag with its
|
||||
# continuation bit set and nothing after it, a length prefix past the end of the payload,
|
||||
# and a fixed32 with two of its four bytes
|
||||
seen = len(switch_states)
|
||||
await raw.send_raw(SWITCH_COMMAND, key + b"\x80")
|
||||
await raw.send_raw(SWITCH_COMMAND, key + tag(2, WIRE_LENGTH) + b"\x7f" + b"ab")
|
||||
await raw.send_raw(SWITCH_COMMAND, tag(1, WIRE_FIXED32) + b"\x01\x02")
|
||||
await raw.send_raw(SWITCH_COMMAND, key + on)
|
||||
await waiter.expect(switch_is(True), label="switch on after truncated frames")
|
||||
assert switch_states[seen:] == [True]
|
||||
await raw.send_raw(SWITCH_COMMAND, key + off)
|
||||
await waiter.expect(switch_is(False))
|
||||
|
||||
# A negative number goes through the fixed32 float path of a normal client
|
||||
client.number_command(number.key, -77.5)
|
||||
await waiter.expect(number_is(-77.5))
|
||||
|
||||
# An unknown field ahead of the known ones is skipped; field 200 needs a two byte tag
|
||||
await raw.send_raw(
|
||||
SWITCH_COMMAND, tag(200, WIRE_VARINT) + encode_varint(300) + key + on
|
||||
)
|
||||
await waiter.expect(switch_is(True))
|
||||
|
||||
# Two byte tags (effect fields 18 and 19) and a two byte varint (300 ms transition)
|
||||
client.light_command(
|
||||
light.key, state=True, brightness=0.5, transition_length=0.3, effect="Pulse"
|
||||
)
|
||||
await waiter.expect(
|
||||
lambda s: (
|
||||
isinstance(s, LightState) and s.key == light.key and s.effect == "Pulse"
|
||||
)
|
||||
)
|
||||
client.light_command(light.key, effect="None", state=False)
|
||||
await waiter.expect(
|
||||
lambda s: isinstance(s, LightState) and s.key == light.key and not s.state
|
||||
)
|
||||
|
||||
# A string whose length prefix needs two varint bytes
|
||||
long_text = "w" * 200
|
||||
client.text_command(text.key, long_text)
|
||||
await waiter.expect(
|
||||
lambda s: (
|
||||
isinstance(s, TextState) and s.key == text.key and s.state == long_text
|
||||
)
|
||||
)
|
||||
@@ -1,37 +0,0 @@
|
||||
"""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])
|
||||
@@ -1,78 +0,0 @@
|
||||
"""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:
|
||||
"""Each superseded field must still be touched by DeviceInfoResponse's
|
||||
generated encode_msg() and calc_size_msg(), i.e. it is still put on the wire.
|
||||
generated encode() and calculate_size(), i.e. it is still put on the wire.
|
||||
"""
|
||||
encode_body = _extract_function_body(CPP_TEXT, "DeviceInfoResponse::encode_msg")
|
||||
size_body = _extract_function_body(CPP_TEXT, "DeviceInfoResponse::calc_size_msg")
|
||||
encode_body = _extract_function_body(CPP_TEXT, "DeviceInfoResponse::encode")
|
||||
size_body = _extract_function_body(CPP_TEXT, "DeviceInfoResponse::calculate_size")
|
||||
for field_name in SUPERSEDED_FIELDS:
|
||||
assert f"msg.{field_name}" in encode_body, (
|
||||
f"DeviceInfoResponse::encode_msg() no longer references {field_name}. "
|
||||
assert f"this->{field_name}" in encode_body, (
|
||||
f"DeviceInfoResponse::encode() no longer references {field_name}. "
|
||||
f"{DEPRECATED_FIELD_TRAP}"
|
||||
)
|
||||
assert f"msg.{field_name}" in size_body, (
|
||||
f"DeviceInfoResponse::calc_size_msg() no longer references "
|
||||
assert f"this->{field_name}" in size_body, (
|
||||
f"DeviceInfoResponse::calculate_size() no longer references "
|
||||
f"{field_name}. {DEPRECATED_FIELD_TRAP}"
|
||||
)
|
||||
|
||||
@@ -380,13 +380,3 @@ def test_api_version_minor_is_at_least_15() -> None:
|
||||
"clients to see api_version >= 1.15 in HelloResponse before they will "
|
||||
"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,13 +15,9 @@ import pytest
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).parents[4] / "script" / "api_protobuf"))
|
||||
|
||||
import aioesphomeapi.api_options_pb2 as pb # noqa: E402
|
||||
from api_protobuf import ( # noqa: E402
|
||||
MAX_MESSAGE_ID,
|
||||
SOURCE_CLIENT,
|
||||
_make_ifdef_line,
|
||||
build_message_type,
|
||||
create_field_type_info,
|
||||
get_varint64_ifdef,
|
||||
validate_message_id,
|
||||
)
|
||||
@@ -38,26 +34,16 @@ def _file_with_messages(
|
||||
file_desc = descriptor_pb2.FileDescriptorProto(name="test.proto")
|
||||
for name, field_type, deprecated in messages:
|
||||
msg = file_desc.message_type.add(name=name)
|
||||
field = msg.field.add()
|
||||
field.CopyFrom(_field(field_type))
|
||||
field = msg.field.add(name="value", number=1, type=field_type)
|
||||
field.options.deprecated = deprecated
|
||||
return file_desc
|
||||
|
||||
|
||||
UINT64 = descriptor_pb2.FieldDescriptorProto.TYPE_UINT64
|
||||
MESSAGE = descriptor_pb2.FieldDescriptorProto.TYPE_MESSAGE
|
||||
DOUBLE = descriptor_pb2.FieldDescriptorProto.TYPE_DOUBLE
|
||||
INT64 = descriptor_pb2.FieldDescriptorProto.TYPE_INT64
|
||||
SINT64 = descriptor_pb2.FieldDescriptorProto.TYPE_SINT64
|
||||
UINT32 = descriptor_pb2.FieldDescriptorProto.TYPE_UINT32
|
||||
INT32 = descriptor_pb2.FieldDescriptorProto.TYPE_INT32
|
||||
SINT32 = descriptor_pb2.FieldDescriptorProto.TYPE_SINT32
|
||||
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:
|
||||
@@ -121,172 +107,3 @@ def test_message_id_at_maximum_is_accepted() -> None:
|
||||
def test_message_id_above_maximum_is_rejected() -> None:
|
||||
with pytest.raises(ValueError, match="exceeds the plaintext"):
|
||||
validate_message_id(MAX_MESSAGE_ID + 1, "TooBigMessage")
|
||||
|
||||
|
||||
def _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
|
||||
|
||||
|
||||
def _decode_case(field_type: int, number: int, *, repeated: bool = False) -> str:
|
||||
"""Return the decode_field() case the generator emits for one decoded field."""
|
||||
field = _field(field_type, number, repeated=repeated)
|
||||
if field_type == MESSAGE:
|
||||
field.type_name = ".Sub"
|
||||
return create_field_type_info(
|
||||
field, needs_decode=True, needs_encode=False
|
||||
).decode_content
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("field_type", "number", "wire_type", "accessor"),
|
||||
[
|
||||
(UINT32, 2, "WIRE_TYPE_VARINT", "value.as_varint()"),
|
||||
(BOOL, 3, "WIRE_TYPE_VARINT", "value.as_bool()"),
|
||||
(STRING, 1, "WIRE_TYPE_LENGTH_DELIMITED", "value.data()"),
|
||||
(FLOAT, 4, "WIRE_TYPE_FIXED32", "value.as_float()"),
|
||||
(FIXED32, 5, "WIRE_TYPE_FIXED32", "value.as_fixed32()"),
|
||||
],
|
||||
)
|
||||
def test_decode_cases_carry_field_number_and_wire_type(
|
||||
field_type: int, number: int, wire_type: str, accessor: str
|
||||
) -> None:
|
||||
"""Each decoded field yields one case keyed on its number and declared wire type."""
|
||||
case = _decode_case(field_type, number)
|
||||
lines = case.splitlines()
|
||||
assert lines[0] == f"case proto_tag({number}, {wire_type}):", case
|
||||
assert accessor in lines[1], case
|
||||
assert lines[-1].strip() == "break;", case
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("field_type", "repeated", "wire_type", "store"),
|
||||
[
|
||||
(UINT32, True, "WIRE_TYPE_VARINT", "this->value.push_back(value.as_varint());"),
|
||||
(
|
||||
STRING,
|
||||
True,
|
||||
"WIRE_TYPE_LENGTH_DELIMITED",
|
||||
"this->value.push_back(value.as_string());",
|
||||
),
|
||||
(
|
||||
MESSAGE,
|
||||
False,
|
||||
"WIRE_TYPE_LENGTH_DELIMITED",
|
||||
"value.decode_to_message(this->value);",
|
||||
),
|
||||
(
|
||||
MESSAGE,
|
||||
True,
|
||||
"WIRE_TYPE_LENGTH_DELIMITED",
|
||||
"value.decode_to_message(this->value.back());",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_repeated_and_message_fields_decode_through_the_same_case_shape(
|
||||
field_type: int, repeated: bool, wire_type: str, store: str
|
||||
) -> None:
|
||||
"""Repeated and sub message fields land in the one switch with their own store."""
|
||||
case = _decode_case(field_type, 7, repeated=repeated)
|
||||
lines = case.splitlines()
|
||||
assert lines[0] == f"case proto_tag(7, {wire_type}):", case
|
||||
assert store in case, case
|
||||
if field_type == MESSAGE and repeated:
|
||||
assert "this->value.emplace_back();" in case, case
|
||||
assert lines[-1].strip() == "break;", case
|
||||
|
||||
|
||||
def test_a_fixed64_field_fails_at_generation_time() -> None:
|
||||
"""The decode loop has no 64 bit wire type path, so such a field must never reach it silently."""
|
||||
desc = descriptor_pb2.DescriptorProto(name="Wide")
|
||||
desc.field.add(name="ratio", number=1, type=DOUBLE)
|
||||
with pytest.raises(
|
||||
ValueError, match="64-bit type 'double' .*ratio.* not supported"
|
||||
):
|
||||
build_message_type(desc, {}, {"Wide": SOURCE_CLIENT})
|
||||
|
||||
|
||||
def test_message_gets_a_single_decode_field_override() -> None:
|
||||
"""All wire types of a decoded message land in one decode_field() switch."""
|
||||
desc = descriptor_pb2.DescriptorProto(name="Mixed")
|
||||
desc.field.add(name="name", number=1, type=STRING)
|
||||
desc.field.add(name="count", number=2, type=UINT32)
|
||||
desc.field.add(name="level", number=3, type=FLOAT)
|
||||
header, cpp, _ = build_message_type(desc, {}, {"Mixed": SOURCE_CLIENT})
|
||||
decl = "void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;"
|
||||
assert header.count(decl) == 1
|
||||
assert (
|
||||
cpp.count(
|
||||
"void Mixed::decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) {"
|
||||
)
|
||||
== 1
|
||||
)
|
||||
assert "switch (tag) {" in cpp
|
||||
assert "const ProtoFieldValue value(data, scalar);" in cpp
|
||||
for number, wire_type in (
|
||||
(1, "WIRE_TYPE_LENGTH_DELIMITED"),
|
||||
(2, "WIRE_TYPE_VARINT"),
|
||||
(3, "WIRE_TYPE_FIXED32"),
|
||||
):
|
||||
assert f"case proto_tag({number}, {wire_type}):" in cpp, cpp
|
||||
|
||||
Reference in New Issue
Block a user