Compare commits

..
71 changed files with 2296 additions and 3097 deletions
+1 -6
View File
@@ -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)
+24 -6
View File
@@ -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;
}
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+134 -198
View File
@@ -287,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);
@@ -319,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
@@ -376,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, (field_id << 3) | 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
@@ -416,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
/**
@@ -692,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
@@ -945,14 +876,19 @@ 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
+2 -1
View File
@@ -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)
-5
View File
@@ -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)
+7 -3
View File
@@ -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)
+2 -1
View File
@@ -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_))
+5 -1
View File
@@ -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]))
+81 -4
View File
@@ -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)
+3 -3
View File
@@ -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);
+29 -11
View File
@@ -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); }
+50 -20
View File
@@ -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(
+2 -1
View File
@@ -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]))
+37
View File
@@ -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
+74 -139
View File
@@ -131,12 +131,6 @@ def force_str(force: bool) -> str:
return str(force).lower()
def _encode_call(func: str, *args: str, force: bool = False) -> str:
"""Emit one ProtoEncode call; every helper takes the cursor and returns it advanced."""
suffix = "_force" if force else ""
return f"pos = ProtoEncode::{func}{suffix}({', '.join(('pos', *args))});"
class TypeInfo(ABC):
"""Base class for all type information."""
@@ -270,16 +264,14 @@ class TypeInfo(ABC):
# 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.
@@ -296,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
@@ -327,43 +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_expr: str) -> str | None:
"""Single-byte tag fixed32 write, or None for multi-byte tags."""
tag = self.calculate_tag()
if tag >= 128:
return None
if self.force:
return _encode_call("write_tag_and_fixed32", str(tag), value_expr)
return (
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 self.fixed32_value_template is not None and (
result := self._encode_fixed32_with_precomputed_tag(
self.fixed32_value_template.format(value=value)
)
):
return result
return _encode_call(self.encode_func, str(self.number), value, force=self.force)
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
@@ -668,8 +635,6 @@ class FloatType(FixedSizeTypeMixin, TypeInfo):
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);"
@@ -732,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
@@ -804,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()
@@ -878,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
@@ -981,9 +951,7 @@ 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:
@@ -1090,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_)"
@@ -1206,13 +1170,9 @@ 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_length_content(self) -> str | None:
@@ -1264,19 +1224,16 @@ 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
@@ -1464,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));"
@@ -1567,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}));"
@@ -1748,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:
@@ -1822,22 +1775,17 @@ class FixedArrayRepeatedType(TypeInfo):
def _encode_element(self, element: str) -> str:
"""Helper to generate encode statement for a single element."""
if isinstance(self._ti, EnumType):
return _encode_call(
self._ti.encode_func,
str(self.number),
f"static_cast<uint32_t>({element})",
force=True,
)
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 _encode_call(
"encode_sub_message", "buffer", str(self.number), element
)
return _encode_call(self._ti.encode_func, str(self.number), element, force=True)
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:
@@ -2189,18 +2137,13 @@ class RepeatedTypeInfo(TypeInfo):
def _encode_element_call(self, element: str) -> str:
"""Helper to generate encode call for a single element."""
if isinstance(self._ti, EnumType):
return _encode_call(
self._ti.encode_func,
str(self.number),
f"static_cast<uint32_t>({element})",
force=True,
)
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 _encode_call(
"encode_sub_message", "buffer", str(self.number), element
)
return _encode_call(self._ti.encode_func, str(self.number), element, force=True)
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:
@@ -2209,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"
@@ -2841,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
@@ -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,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');
-40
View File
@@ -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,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,11 +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,
_make_ifdef_line,
create_field_type_info,
get_varint64_ifdef,
validate_message_id,
)
@@ -45,14 +43,7 @@ UINT64 = descriptor_pb2.FieldDescriptorProto.TYPE_UINT64
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:
@@ -116,69 +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