[api] Collapse the three protobuf decode virtuals into one (#19016)

This commit is contained in:
J. Nick Koston
2026-09-24 15:30:11 +01:00
committed by GitHub
parent 7dc98c683b
commit f1f83cda07
10 changed files with 1351 additions and 1730 deletions
File diff suppressed because it is too large Load Diff
+60 -107
View File
@@ -429,8 +429,7 @@ class HelloRequest final : public ProtoDecodableMessage {
#endif
protected:
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class HelloResponse final : public ProtoMessage {
public:
@@ -474,7 +473,7 @@ class DisconnectRequest final : public ProtoDecodableMessage {
#endif
protected:
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class DisconnectResponse final : public ProtoMessage {
public:
@@ -850,8 +849,7 @@ class CoverCommandRequest final : public CommandProtoMessage {
#endif
protected:
bool decode_32bit(uint32_t field_id, Proto32Bit value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
#endif
#ifdef USE_FAN
@@ -925,9 +923,7 @@ class FanCommandRequest final : public CommandProtoMessage {
#endif
protected:
bool decode_32bit(uint32_t field_id, Proto32Bit value) override;
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
#endif
#ifdef USE_LIGHT
@@ -1023,9 +1019,7 @@ class LightCommandRequest final : public CommandProtoMessage {
#endif
protected:
bool decode_32bit(uint32_t field_id, Proto32Bit value) override;
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
#endif
#ifdef USE_SENSOR
@@ -1130,8 +1124,7 @@ class SwitchCommandRequest final : public CommandProtoMessage {
#endif
protected:
bool decode_32bit(uint32_t field_id, Proto32Bit value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
#endif
#ifdef USE_TEXT_SENSOR
@@ -1191,7 +1184,7 @@ class SubscribeLogsRequest final : public ProtoDecodableMessage {
#endif
protected:
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class SubscribeLogsResponse final : public ProtoMessage {
public:
@@ -1234,7 +1227,7 @@ class NoiseEncryptionSetKeyRequest final : public ProtoDecodableMessage {
#endif
protected:
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class NoiseEncryptionSetKeyResponse final : public ProtoMessage {
public:
@@ -1328,8 +1321,7 @@ class HomeassistantActionResponse final : public ProtoDecodableMessage {
#endif
protected:
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
#endif
#ifdef USE_API_HOMEASSISTANT_STATES
@@ -1370,7 +1362,7 @@ class HomeAssistantStateResponse final : public ProtoDecodableMessage {
#endif
protected:
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
#endif
class GetTimeRequest final : public ProtoMessage {
@@ -1399,7 +1391,7 @@ class DSTRule final : public ProtoDecodableMessage {
#endif
protected:
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class ParsedTimezone final : public ProtoDecodableMessage {
public:
@@ -1412,8 +1404,7 @@ class ParsedTimezone final : public ProtoDecodableMessage {
#endif
protected:
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class GetTimeResponse final : public ProtoDecodableMessage {
public:
@@ -1430,8 +1421,7 @@ class GetTimeResponse final : public ProtoDecodableMessage {
#endif
protected:
bool decode_32bit(uint32_t field_id, Proto32Bit value) override;
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
#ifdef USE_API_USER_DEFINED_ACTIONS
class ListEntitiesServicesArgument final : public ProtoMessage {
@@ -1499,9 +1489,7 @@ class ExecuteServiceArgument final : public ProtoDecodableMessage {
#endif
protected:
bool decode_32bit(uint32_t field_id, Proto32Bit value) override;
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class ExecuteServiceRequest final : public ProtoDecodableMessage {
public:
@@ -1524,9 +1512,7 @@ class ExecuteServiceRequest final : public ProtoDecodableMessage {
#endif
protected:
bool decode_32bit(uint32_t field_id, Proto32Bit value) override;
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
#endif
#ifdef USE_API_USER_DEFINED_ACTION_RESPONSES
@@ -1617,7 +1603,7 @@ class CameraImageRequest final : public ProtoDecodableMessage {
#endif
protected:
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
#endif
#ifdef USE_CLIMATE
@@ -1723,9 +1709,7 @@ class ClimateCommandRequest final : public CommandProtoMessage {
#endif
protected:
bool decode_32bit(uint32_t field_id, Proto32Bit value) override;
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
#endif
#ifdef USE_WATER_HEATER
@@ -1797,8 +1781,7 @@ class WaterHeaterCommandRequest final : public CommandProtoMessage {
#endif
protected:
bool decode_32bit(uint32_t field_id, Proto32Bit value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
#endif
#ifdef USE_NUMBER
@@ -1861,8 +1844,7 @@ class NumberCommandRequest final : public CommandProtoMessage {
#endif
protected:
bool decode_32bit(uint32_t field_id, Proto32Bit value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
#endif
#ifdef USE_SELECT
@@ -1920,9 +1902,7 @@ class SelectCommandRequest final : public CommandProtoMessage {
#endif
protected:
bool decode_32bit(uint32_t field_id, Proto32Bit value) override;
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
#endif
#ifdef USE_SIREN
@@ -1988,9 +1968,7 @@ class SirenCommandRequest final : public CommandProtoMessage {
#endif
protected:
bool decode_32bit(uint32_t field_id, Proto32Bit value) override;
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
#endif
#ifdef USE_LOCK
@@ -2052,9 +2030,7 @@ class LockCommandRequest final : public CommandProtoMessage {
#endif
protected:
bool decode_32bit(uint32_t field_id, Proto32Bit value) override;
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
#endif
#ifdef USE_BUTTON
@@ -2090,8 +2066,7 @@ class ButtonCommandRequest final : public CommandProtoMessage {
#endif
protected:
bool decode_32bit(uint32_t field_id, Proto32Bit value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
#endif
#ifdef USE_MEDIA_PLAYER
@@ -2177,9 +2152,7 @@ class MediaPlayerCommandRequest final : public CommandProtoMessage {
#endif
protected:
bool decode_32bit(uint32_t field_id, Proto32Bit value) override;
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
#endif
#ifdef USE_BLUETOOTH_PROXY
@@ -2196,7 +2169,7 @@ class SubscribeBluetoothLEAdvertisementsRequest final : public ProtoDecodableMes
#endif
protected:
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class BluetoothLERawAdvertisement final : public ProtoMessage {
public:
@@ -2250,7 +2223,7 @@ class BluetoothDeviceRequest final : public ProtoDecodableMessage {
#endif
protected:
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class BluetoothDeviceConnectionResponse final : public ProtoMessage {
public:
@@ -2288,7 +2261,7 @@ class BluetoothGATTGetServicesRequest final : public ProtoDecodableMessage {
#endif
protected:
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class BluetoothGATTDescriptor final : public ProtoMessage {
public:
@@ -2399,7 +2372,7 @@ class BluetoothGATTReadRequest final : public ProtoDecodableMessage {
#endif
protected:
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class BluetoothGATTReadResponse final : public ProtoMessage {
public:
@@ -2445,8 +2418,7 @@ class BluetoothGATTWriteRequest final : public ProtoDecodableMessage {
#endif
protected:
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class BluetoothGATTReadDescriptorRequest final : public ProtoDecodableMessage {
public:
@@ -2462,7 +2434,7 @@ class BluetoothGATTReadDescriptorRequest final : public ProtoDecodableMessage {
#endif
protected:
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class BluetoothGATTWriteDescriptorRequest final : public ProtoDecodableMessage {
public:
@@ -2480,8 +2452,7 @@ class BluetoothGATTWriteDescriptorRequest final : public ProtoDecodableMessage {
#endif
protected:
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class BluetoothGATTNotifyRequest final : public ProtoDecodableMessage {
public:
@@ -2498,7 +2469,7 @@ class BluetoothGATTNotifyRequest final : public ProtoDecodableMessage {
#endif
protected:
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class BluetoothGATTNotifyDataResponse final : public ProtoMessage {
public:
@@ -2716,7 +2687,7 @@ class BluetoothScannerSetModeRequest final : public ProtoDecodableMessage {
#endif
protected:
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
#endif
#ifdef USE_VOICE_ASSISTANT
@@ -2734,7 +2705,7 @@ class SubscribeVoiceAssistantRequest final : public ProtoDecodableMessage {
#endif
protected:
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class VoiceAssistantAudioSettings final : public ProtoMessage {
public:
@@ -2791,7 +2762,7 @@ class VoiceAssistantResponse final : public ProtoDecodableMessage {
#endif
protected:
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class VoiceAssistantEventData final : public ProtoDecodableMessage {
public:
@@ -2802,7 +2773,7 @@ class VoiceAssistantEventData final : public ProtoDecodableMessage {
#endif
protected:
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class VoiceAssistantEventResponse final : public ProtoDecodableMessage {
public:
@@ -2818,8 +2789,7 @@ class VoiceAssistantEventResponse final : public ProtoDecodableMessage {
#endif
protected:
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class VoiceAssistantAudio final : public ProtoDecodableMessage {
public:
@@ -2844,8 +2814,7 @@ class VoiceAssistantAudio final : public ProtoDecodableMessage {
#endif
protected:
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class VoiceAssistantTimerEventResponse final : public ProtoDecodableMessage {
public:
@@ -2865,8 +2834,7 @@ class VoiceAssistantTimerEventResponse final : public ProtoDecodableMessage {
#endif
protected:
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class VoiceAssistantAnnounceRequest final : public ProtoDecodableMessage {
public:
@@ -2884,8 +2852,7 @@ class VoiceAssistantAnnounceRequest final : public ProtoDecodableMessage {
#endif
protected:
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class VoiceAssistantAnnounceFinished final : public ProtoMessage {
public:
@@ -2938,8 +2905,7 @@ class VoiceAssistantExternalWakeWord final : public ProtoDecodableMessage {
#endif
protected:
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class VoiceAssistantConfigurationRequest final : public ProtoDecodableMessage {
public:
@@ -2954,7 +2920,7 @@ class VoiceAssistantConfigurationRequest final : public ProtoDecodableMessage {
#endif
protected:
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class VoiceAssistantConfigurationResponse final : public ProtoMessage {
public:
@@ -2991,7 +2957,7 @@ class VoiceAssistantSetConfiguration final : public ProtoDecodableMessage {
#endif
protected:
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
#endif
#ifdef USE_ALARM_CONTROL_PANEL
@@ -3051,9 +3017,7 @@ class AlarmControlPanelCommandRequest final : public CommandProtoMessage {
#endif
protected:
bool decode_32bit(uint32_t field_id, Proto32Bit value) override;
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
#endif
#ifdef USE_TEXT
@@ -3114,9 +3078,7 @@ class TextCommandRequest final : public CommandProtoMessage {
#endif
protected:
bool decode_32bit(uint32_t field_id, Proto32Bit value) override;
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
#endif
#ifdef USE_DATETIME_DATE
@@ -3177,8 +3139,7 @@ class DateCommandRequest final : public CommandProtoMessage {
#endif
protected:
bool decode_32bit(uint32_t field_id, Proto32Bit value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
#endif
#ifdef USE_DATETIME_TIME
@@ -3239,8 +3200,7 @@ class TimeCommandRequest final : public CommandProtoMessage {
#endif
protected:
bool decode_32bit(uint32_t field_id, Proto32Bit value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
#endif
#ifdef USE_EVENT
@@ -3346,8 +3306,7 @@ class ValveCommandRequest final : public CommandProtoMessage {
#endif
protected:
bool decode_32bit(uint32_t field_id, Proto32Bit value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
#endif
#ifdef USE_DATETIME_DATETIME
@@ -3404,8 +3363,7 @@ class DateTimeCommandRequest final : public CommandProtoMessage {
#endif
protected:
bool decode_32bit(uint32_t field_id, Proto32Bit value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
#endif
#ifdef USE_UPDATE
@@ -3470,8 +3428,7 @@ class UpdateCommandRequest final : public CommandProtoMessage {
#endif
protected:
bool decode_32bit(uint32_t field_id, Proto32Bit value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
#endif
#ifdef USE_ZWAVE_PROXY
@@ -3495,7 +3452,7 @@ class ZWaveProxyFrame final : public ProtoDecodableMessage {
#endif
protected:
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class ZWaveProxyRequest final : public ProtoDecodableMessage {
public:
@@ -3518,8 +3475,7 @@ class ZWaveProxyRequest final : public ProtoDecodableMessage {
#endif
protected:
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class ZWaveProxyRequestResponse final : public ProtoMessage {
public:
@@ -3589,9 +3545,7 @@ class InfraredRFTransmitRawTimingsRequest final : public ProtoDecodableMessage {
#endif
protected:
bool decode_32bit(uint32_t field_id, Proto32Bit value) override;
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class InfraredRFReceiveEvent final : public ProtoMessage {
public:
@@ -3662,7 +3616,7 @@ class SerialProxyConfigureRequest final : public ProtoDecodableMessage {
#endif
protected:
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class SerialProxyDataReceived final : public ProtoMessage {
public:
@@ -3705,8 +3659,7 @@ class SerialProxyWriteRequest final : public ProtoDecodableMessage {
#endif
protected:
bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class SerialProxySetModemPinsRequest final : public ProtoDecodableMessage {
public:
@@ -3722,7 +3675,7 @@ class SerialProxySetModemPinsRequest final : public ProtoDecodableMessage {
#endif
protected:
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class SerialProxyGetModemPinsRequest final : public ProtoDecodableMessage {
public:
@@ -3737,7 +3690,7 @@ class SerialProxyGetModemPinsRequest final : public ProtoDecodableMessage {
#endif
protected:
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class SerialProxyGetModemPinsResponse final : public ProtoMessage {
public:
@@ -3775,7 +3728,7 @@ class SerialProxyRequest final : public ProtoDecodableMessage {
#endif
protected:
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class SerialProxyRequestResponse final : public ProtoMessage {
public:
@@ -3814,7 +3767,7 @@ class SerialProxySetModeRequest final : public ProtoDecodableMessage {
#endif
protected:
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
#endif
#ifdef USE_BLUETOOTH_PROXY_CONNECTIONS
@@ -3835,7 +3788,7 @@ class BluetoothSetConnectionParamsRequest final : public ProtoDecodableMessage {
#endif
protected:
bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;
void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;
};
class BluetoothSetConnectionParamsResponse final : public ProtoMessage {
public:
+59 -58
View File
@@ -220,73 +220,74 @@ void ProtoDecodableMessage::decode(const uint8_t *buffer, size_t length) {
const uint8_t *ptr = buffer;
const uint8_t *end = buffer + length;
while (ptr < end) {
// Parse field header - ptr < end guarantees len >= 1
// Single-byte varints dominate, so that case advances the cursor inline.
auto read_varint = [&](proto_varint_value_t &value) ESPHOME_ALWAYS_INLINE {
if (ptr == end)
return false;
if (*ptr < 0x80) [[likely]] {
value = *ptr++;
return true;
}
auto res = ProtoVarInt::parse_non_empty(ptr, end - ptr);
if (!res.has_value()) {
if (!res.has_value())
return false;
value = res.value;
ptr += res.consumed;
return true;
};
while (ptr < end) {
proto_varint_value_t tag_value;
if (!read_varint(tag_value)) {
ESP_LOGV(TAG, "Invalid field start at offset %ld", (long) (ptr - buffer));
return;
}
uint32_t tag = static_cast<uint32_t>(res.value);
uint32_t tag = static_cast<uint32_t>(tag_value);
uint32_t field_type = tag & WIRE_TYPE_MASK;
uint32_t field_id = tag >> 3;
ptr += res.consumed;
// Length-delimited fields move this past the length prefix
const uint8_t *data = ptr;
proto_varint_value_t scalar;
switch (field_type) {
case WIRE_TYPE_VARINT: { // VarInt
res = ProtoVarInt::parse(ptr, end - ptr);
if (!res.has_value()) {
ESP_LOGV(TAG, "Invalid VarInt at offset %ld", (long) (ptr - buffer));
return;
}
if (!this->decode_varint(field_id, res.value)) {
ESP_LOGV(TAG, "Cannot decode VarInt field %" PRIu32 " with value %" PRIu64 "!", field_id,
static_cast<uint64_t>(res.value));
}
ptr += res.consumed;
break;
}
case WIRE_TYPE_LENGTH_DELIMITED: { // Length-delimited
res = ProtoVarInt::parse(ptr, end - ptr);
if (!res.has_value()) {
ESP_LOGV(TAG, "Invalid Length Delimited at offset %ld", (long) (ptr - buffer));
return;
}
uint32_t field_length = static_cast<uint32_t>(res.value);
ptr += res.consumed;
if (field_length > static_cast<size_t>(end - ptr)) {
ESP_LOGV(TAG, "Out-of-bounds Length Delimited at offset %ld", (long) (ptr - buffer));
return;
}
if (!this->decode_length(field_id, ProtoLengthDelimited(ptr, field_length))) {
ESP_LOGV(TAG, "Cannot decode Length Delimited field %" PRIu32 "!", field_id);
}
ptr += field_length;
break;
}
case WIRE_TYPE_FIXED32: { // 32-bit
if (end - ptr < 4) {
ESP_LOGV(TAG, "Out-of-bounds Fixed32-bit at offset %ld", (long) (ptr - buffer));
return;
}
uint32_t val;
#if __BYTE_ORDER__ == __ORDER_LITTLE_ENDIAN__
// Protobuf fixed32 is little-endian — direct load on LE platforms
memcpy(&val, ptr, 4);
#else
val = encode_uint32(ptr[3], ptr[2], ptr[1], ptr[0]);
#endif
if (!this->decode_32bit(field_id, Proto32Bit(val))) {
ESP_LOGV(TAG, "Cannot decode 32-bit field %" PRIu32 " with value %" PRIu32 "!", field_id, val);
}
ptr += 4;
break;
}
default:
ESP_LOGV(TAG, "Invalid field type %" PRIu32 " at offset %ld", field_type, (long) (ptr - buffer));
if (field_type == WIRE_TYPE_VARINT) [[likely]] {
if (!read_varint(scalar)) {
ESP_LOGV(TAG, "Invalid VarInt at offset %ld", (long) (ptr - buffer));
return;
}
} else {
switch (field_type) {
case WIRE_TYPE_LENGTH_DELIMITED: {
proto_varint_value_t length_value;
if (!read_varint(length_value)) {
ESP_LOGV(TAG, "Invalid Length Delimited at offset %ld", (long) (ptr - buffer));
return;
}
uint32_t field_length = static_cast<uint32_t>(length_value);
if (field_length > static_cast<size_t>(end - ptr)) {
ESP_LOGV(TAG, "Out-of-bounds Length Delimited at offset %ld", (long) (ptr - buffer));
return;
}
data = ptr;
scalar = field_length;
ptr += field_length;
break;
}
case WIRE_TYPE_FIXED32: {
if (end - ptr < 4) {
ESP_LOGV(TAG, "Out-of-bounds Fixed32-bit at offset %ld", (long) (ptr - buffer));
return;
}
// Byte loads instead of memcpy: ESP-IDF passes -fno-builtin-memcpy, which made this a call
scalar = encode_uint32(ptr[3], ptr[2], ptr[1], ptr[0]);
ptr += 4;
break;
}
default:
ESP_LOGV(TAG, "Invalid field type %" PRIu32 " at offset %ld", field_type, (long) (ptr - buffer));
return;
}
}
this->decode_field(tag, data, scalar);
}
}
+31 -33
View File
@@ -170,40 +170,43 @@ class ProtoVarInt {
class ProtoMessage;
class ProtoSize;
class ProtoLengthDelimited {
/// Case label for decode_field(): the wire tag of a field, so a field that arrives with another wire
/// type matches no case.
constexpr uint32_t proto_tag(uint32_t field_id, uint32_t wire_type) { return (field_id << 3) | wire_type; }
/// One decoded field: the payload pointer and a scalar holding the varint or fixed32 value, or the
/// length of a length-delimited field. The wire type in the tag says which applies; accessors do not check.
class ProtoFieldValue {
public:
explicit ProtoLengthDelimited(const uint8_t *value, size_t length) : value_(value), length_(length) {}
std::string as_string() const { return std::string(reinterpret_cast<const char *>(this->value_), this->length_); }
ProtoFieldValue(const uint8_t *data, proto_varint_value_t scalar) : data_(data), scalar_(scalar) {}
// Direct access to raw data without string allocation
const uint8_t *data() const { return this->value_; }
size_t size() const { return this->length_; }
proto_varint_value_t as_varint() const { return this->scalar_; }
// A bool is sent as 0 or 1, so the low word is enough and saves a second compare with 64 bit varints
bool as_bool() const { return static_cast<uint32_t>(this->scalar_) != 0; }
/// Decode the length-delimited data into a message instance.
// Length-delimited accessors
const uint8_t *data() const { return this->data_; }
size_t size() const { return static_cast<size_t>(this->scalar_); }
std::string as_string() const { return std::string(reinterpret_cast<const char *>(this->data_), this->size()); }
/// Decode the length-delimited payload into a message instance.
/// Template preserves concrete type so decode() resolves statically.
template<typename T> void decode_to_message(T &msg) const;
template<typename T> void decode_to_message(T &msg) const { msg.decode(this->data_, this->size()); }
protected:
const uint8_t *const value_;
const size_t length_;
};
class Proto32Bit {
public:
explicit Proto32Bit(uint32_t value) : value_(value) {}
uint32_t as_fixed32() const { return this->value_; }
int32_t as_sfixed32() const { return static_cast<int32_t>(this->value_); }
// Fixed32 accessors
uint32_t as_fixed32() const { return static_cast<uint32_t>(this->scalar_); }
int32_t as_sfixed32() const { return static_cast<int32_t>(this->as_fixed32()); }
float as_float() const {
union {
uint32_t raw;
float value;
} s{};
s.raw = this->value_;
s.raw = this->as_fixed32();
return s.value;
}
protected:
const uint32_t value_;
private:
const uint8_t *data_;
proto_varint_value_t scalar_;
};
// NOTE: Proto64Bit class removed - wire type 1 (64-bit fixed) not supported
@@ -255,7 +258,7 @@ class ProtoWriteBuffer {
*
* Following https://protobuf.dev/programming-guides/encoding/#structure
*/
void encode_field_raw(uint32_t field_id, uint32_t type) { this->encode_varint_raw((field_id << 3) | type); }
void encode_field_raw(uint32_t field_id, uint32_t type) { this->encode_varint_raw(proto_tag(field_id, type)); }
/// Single-pass encode for repeated submessage elements.
/// Thin template wrapper; all buffer work is in the non-template core.
template<typename T> void encode_sub_message(uint32_t field_id, const T &value);
@@ -385,7 +388,7 @@ class ProtoEncode {
}
[[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);
return encode_varint_raw(pos PROTO_ENCODE_DEBUG_ARG, proto_tag(field_id, type));
}
/// Write a single precomputed tag byte. Tag must be < 128.
[[nodiscard]] static inline uint8_t *ESPHOME_ALWAYS_INLINE
@@ -735,10 +738,10 @@ class ProtoDecodableMessage : public ProtoMessage {
protected:
~ProtoDecodableMessage() = default;
virtual bool decode_varint(uint32_t field_id, proto_varint_value_t value) { return false; }
virtual bool decode_length(uint32_t field_id, ProtoLengthDelimited value) { return false; }
virtual bool decode_32bit(uint32_t field_id, Proto32Bit value) { return false; }
// NOTE: decode_64bit removed - wire type 1 not supported
/// Store one decoded field; \p scalar is the varint or fixed32 value, or the length of the
/// length-delimited payload at \p data. An unknown field or wrong wire type matches no case and is skipped.
/// Three register arguments keep the decode loop free of spills.
virtual void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) {}
};
class ProtoSize {
@@ -864,7 +867,7 @@ class ProtoSize {
* @return The number of bytes needed to encode the field ID and wire type
*/
static constexpr uint32_t field(uint32_t field_id, uint32_t type) {
uint32_t tag = (field_id << 3) | (type & WIRE_TYPE_MASK);
uint32_t tag = proto_tag(field_id, type & WIRE_TYPE_MASK);
return varint(tag);
}
@@ -958,11 +961,6 @@ template<typename T> inline void ProtoWriteBuffer::encode_optional_sub_message(u
this->encode_optional_sub_message(field_id, T::calc_size_msg(&value), &value, &T::encode_msg);
}
// Template decode_to_message - preserves concrete type so decode() resolves statically
template<typename T> void ProtoLengthDelimited::decode_to_message(T &msg) const {
msg.decode(this->value_, this->length_);
}
template<typename T> const char *proto_enum_to_string(T value);
// ProtoService removed — its methods were inlined into APIConnection.
+103 -242
View File
@@ -28,6 +28,11 @@ class WireType(IntEnum):
END_GROUP = 4 # groups (deprecated)
FIXED32 = 5 # fixed32, sfixed32, float
@property
def cpp_name(self) -> str:
"""The matching constant in proto.h."""
return f"WIRE_TYPE_{self.name}"
# Generate with
# protoc --python_out=script/api_protobuf -I esphome/components/api/ api_options.proto
@@ -224,41 +229,23 @@ class TypeInfo(ABC):
def class_member(self) -> str:
return f"{self.cpp_type} {self.field_name}{{{self.default_value}}};"
@property
def decode_varint_content(self) -> str:
content = self.decode_varint
if content is None:
return None
return f"case {self.number}: this->{self.field_name} = {content}; break;"
def decode_case(self, body: str) -> str:
"""Emit one decode_field() case, keyed on the field's wire tag."""
return f"case proto_tag({self.number}, {self.wire_type.cpp_name}):\n" + indent(
f"{body}\nbreak;"
)
decode_varint = None
# Expression that reads this field from `value`; None when the type is never decoded.
decode_expr: str | None = None
def _decode_store(self, expr: str) -> str:
return f"this->{self.field_name} = {expr};"
@property
def decode_length_content(self) -> str:
content = self.decode_length
if content is None:
return None
return f"case {self.number}: this->{self.field_name} = {content}; break;"
decode_length = None
@property
def decode_32bit_content(self) -> str:
content = self.decode_32bit
if content is None:
return None
return f"case {self.number}: this->{self.field_name} = {content}; break;"
decode_32bit = None
@property
def decode_64bit_content(self) -> str:
content = self.decode_64bit
if content is None:
return None
return f"case {self.number}: this->{self.field_name} = {content}; break;"
decode_64bit = None
def decode_content(self) -> str | None:
"""The decode_field() case for this field, or None when it is never decoded."""
expr = self.decode_expr
return None if expr is None else self.decode_case(self._decode_store(expr))
# Mapping from encode_func to raw encode expression template.
# When a forced field has a single-byte tag, the code generator emits
@@ -334,11 +321,12 @@ class TypeInfo(ABC):
)
)
def _encode_fixed32_with_precomputed_tag(self, value_expr: str) -> str | None:
"""Single-byte tag fixed32 write, or None for multi-byte tags."""
def _encode_fixed32_with_precomputed_tag(self, value: str) -> str | None:
"""Single-byte tag fixed32 write, or None for other types and multi-byte tags."""
tag = self.calculate_tag()
if tag >= 128:
if self.fixed32_value_template is None or tag >= 128:
return None
value_expr = self.fixed32_value_template.format(value=value)
if self.force:
return _encode_call("write_tag_and_fixed32", str(tag), value_expr)
return (
@@ -352,14 +340,14 @@ class TypeInfo(ABC):
value = f"this->{self.field_name}"
if result := self._encode_with_precomputed_tag(value):
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)
)
):
if result := self._encode_fixed32_with_precomputed_tag(value):
return result
return _encode_call(self.encode_func, str(self.number), value, force=self.force)
def encode_element(self, number: int, element: str) -> str:
"""Encode one element of a repeated field; elements are always written."""
return _encode_call(self.encode_func, str(number), element, force=True)
encode_func = None
@classmethod
@@ -633,7 +621,6 @@ class DoubleType(FixedSizeTypeMixin, TypeInfo):
# Unsupported but defined for completeness
cpp_type = "double"
default_value = "0.0"
decode_64bit = "value.as_double()"
encode_func = "encode_double"
wire_type = WireType.FIXED64 # Uses wire type 1 according to protobuf spec
@@ -659,7 +646,7 @@ class DoubleType(FixedSizeTypeMixin, TypeInfo):
class FloatType(FixedSizeTypeMixin, TypeInfo):
cpp_type = "float"
default_value = "0.0f"
decode_32bit = "value.as_float()"
decode_expr = "value.as_float()"
encode_func = "encode_float"
wire_type = WireType.FIXED32 # Uses wire type 5
@@ -688,7 +675,7 @@ class Int64Type(VarintTypeMixin, TypeInfo):
cpp_type = "int64_t"
_varint_max_bits = 64
default_value = "0"
decode_varint = "static_cast<int64_t>(value)"
decode_expr = "static_cast<int64_t>(value.as_varint())"
encode_func = "encode_int64"
wire_type = WireType.VARINT # Uses wire type 0
@@ -709,7 +696,7 @@ class UInt64Type(VarintTypeMixin, TypeInfo):
cpp_type = "uint64_t"
_varint_max_bits = 64
default_value = "0"
decode_varint = "value"
decode_expr = "value.as_varint()"
encode_func = "encode_uint64"
wire_type = WireType.VARINT # Uses wire type 0
@@ -744,7 +731,7 @@ class Int32Type(VarintTypeMixin, TypeInfo):
cpp_type = "int32_t"
_varint_max_bits = 64 # int32 is sign-extended to 64 bits in protobuf
default_value = "0"
decode_varint = "static_cast<int32_t>(value)"
decode_expr = "static_cast<int32_t>(value.as_varint())"
encode_func = "encode_int32"
wire_type = WireType.VARINT # Uses wire type 0
@@ -764,7 +751,6 @@ class Int32Type(VarintTypeMixin, TypeInfo):
class Fixed64Type(FixedSizeTypeMixin, TypeInfo):
cpp_type = "uint64_t"
default_value = "0"
decode_64bit = "value.as_fixed64()"
encode_func = "encode_fixed64"
wire_type = WireType.FIXED64 # Uses wire type 1
@@ -790,7 +776,7 @@ class Fixed64Type(FixedSizeTypeMixin, TypeInfo):
class Fixed32Type(FixedSizeTypeMixin, TypeInfo):
cpp_type = "uint32_t"
default_value = "0"
decode_32bit = "value.as_fixed32()"
decode_expr = "value.as_fixed32()"
encode_func = "encode_fixed32"
wire_type = WireType.FIXED32 # Uses wire type 5
@@ -819,7 +805,7 @@ class BoolType(VarintTypeMixin, TypeInfo):
_varint_max_bits = 1
cpp_type = "bool"
default_value = "false"
decode_varint = "value != 0"
decode_expr = "value.as_bool()"
encode_func = "encode_bool"
wire_type = WireType.VARINT # Uses wire type 0
@@ -839,7 +825,7 @@ class StringType(TypeInfo):
default_value = ""
reference_type = "std::string &"
const_reference_type = "const std::string &"
decode_length = "value.as_string()"
decode_expr = "value.as_string()"
encode_func = "encode_string"
wire_type = WireType.LENGTH_DELIMITED # Uses wire type 2
@@ -954,6 +940,9 @@ class MessageType(TypeInfo):
def can_use_dump_field(cls) -> bool:
return False
def encode_element(self, number: int, element: str) -> str:
return _encode_call("encode_sub_message", "buffer", str(number), element)
@property
def cpp_type(self) -> str:
return self._field.type_name[1:]
@@ -980,14 +969,6 @@ class MessageType(TypeInfo):
self.encode_func, "buffer", str(self.number), f"this->{self.field_name}"
)
@property
def decode_length(self) -> str:
# Override to return None for message types because we can't use template-based
# decoding when the specific message type isn't known at compile time.
# Instead, we use the non-template decode_to_message() method which allows
# runtime polymorphism through virtual function calls.
return None
@property
def public_content(self) -> list[str]:
content = [self.class_member]
@@ -1003,19 +984,14 @@ class MessageType(TypeInfo):
)
@property
def decode_length_content(self) -> str:
# Custom decode that doesn't use templates
def decode_content(self) -> str:
body = f"value.decode_to_message(this->{self.field_name});"
if self._track_presence:
# decode_to_message() cannot report failure, so setting the flag
# afterwards only documents intent; a status-returning decode could
# gate it for real without touching callers.
return (
f"case {self.number}:\n"
f" value.decode_to_message(this->{self.field_name});\n"
f" this->has_{self.name} = true;\n"
f" break;"
)
return f"case {self.number}: value.decode_to_message(this->{self.field_name}); break;"
body += f"\nthis->has_{self.name} = true;"
return self.decode_case(body)
def dump(self, name: str) -> str:
return f"{name}.dump_to(out);"
@@ -1054,7 +1030,7 @@ class BytesType(TypeInfo):
reference_type = "std::string &"
const_reference_type = "const std::string &"
encode_func = "encode_bytes"
decode_length = "value.as_string()"
decode_expr = "value.as_string()"
wire_type = WireType.LENGTH_DELIMITED # Uses wire type 2
@property
@@ -1164,11 +1140,6 @@ class PointerToBufferTypeBase(TypeInfo):
super().__init__(field)
self.array_size = 0
@property
def decode_length(self) -> str | None:
# This is handled in decode_length_content
return None
@property
def wire_type(self) -> WireType:
"""Get the wire type for this field."""
@@ -1210,12 +1181,11 @@ class PointerToBytesBufferType(PointerToBufferTypeBase):
)
@property
def decode_length_content(self) -> str | None:
return f"""case {self.number}: {{
this->{self.field_name} = value.data();
this->{self.field_name}_len = value.size();
break;
}}"""
def decode_content(self) -> str:
return self.decode_case(
f"this->{self.field_name} = value.data();\n"
f"this->{self.field_name}_len = value.size();",
)
def dump(self, name: str) -> str:
return (
@@ -1275,11 +1245,10 @@ class PointerToStringBufferType(PointerToBufferTypeBase):
)
@property
def decode_length_content(self) -> str | None:
return f"""case {self.number}: {{
this->{self.field_name} = StringRef(reinterpret_cast<const char *>(value.data()), value.size());
break;
}}"""
def decode_content(self) -> str:
return self.decode_case(
f"this->{self.field_name} = StringRef(value.data(), value.size());",
)
def dump(self, name: str) -> str:
# Not used since we use dump_field, but required by abstract base class
@@ -1348,14 +1317,13 @@ class PackedBufferTypeInfo(TypeInfo):
]
@property
def decode_length_content(self) -> str:
def decode_content(self) -> str:
"""Store pointer to buffer and calculate count of packed varints."""
return f"""case {self.number}: {{
this->{self.field_name}_data_ = value.data();
this->{self.field_name}_length_ = value.size();
this->{self.field_name}_count_ = count_packed_varints(value.data(), value.size());
break;
}}"""
return self.decode_case(
f"this->{self.field_name}_data_ = value.data();\n"
f"this->{self.field_name}_length_ = value.size();\n"
f"this->{self.field_name}_count_ = count_packed_varints(value.data(), value.size());",
)
@property
def encode_content(self) -> str:
@@ -1440,17 +1408,11 @@ class FixedArrayBytesType(TypeInfo):
]
@property
def decode_length_content(self) -> str:
o = f"case {self.number}: {{\n"
o += " const std::string &data_str = value.as_string();\n"
o += f" this->{self.field_name}_len = data_str.size();\n"
o += f" if (this->{self.field_name}_len > {self.array_size}) {{\n"
o += f" this->{self.field_name}_len = {self.array_size};\n"
o += " }\n"
o += f" memcpy(this->{self.field_name}, data_str.data(), this->{self.field_name}_len);\n"
o += " break;\n"
o += "}"
return o
def decode_content(self) -> str:
return self.decode_case(
f"this->{self.field_name}_len = std::min<size_t>(value.size(), {self.array_size});\n"
f"memcpy(this->{self.field_name}, value.data(), this->{self.field_name}_len);",
)
@property
def encode_content(self) -> str:
@@ -1513,7 +1475,7 @@ class UInt32Type(VarintTypeMixin, TypeInfo):
cpp_type = "uint32_t"
_varint_max_bits = 32
default_value = "0"
decode_varint = "value"
decode_expr = "value.as_varint()"
encode_func = "encode_uint32"
wire_type = WireType.VARINT # Uses wire type 0
@@ -1536,13 +1498,21 @@ class UInt32Type(VarintTypeMixin, TypeInfo):
class EnumType(VarintTypeMixin, TypeInfo):
_varint_max_bits = 32
def encode_element(self, number: int, element: str) -> str:
return _encode_call(
self.encode_func,
str(number),
f"static_cast<uint32_t>({element})",
force=True,
)
@property
def cpp_type(self) -> str:
return f"enums::{self._field.type_name[1:]}"
@property
def decode_varint(self) -> str:
return f"static_cast<{self.cpp_type}>(value)"
def decode_expr(self) -> str:
return f"static_cast<{self.cpp_type}>(value.as_varint())"
default_value = ""
wire_type = WireType.VARINT # Uses wire type 0
@@ -1589,7 +1559,7 @@ class EnumType(VarintTypeMixin, TypeInfo):
class SFixed32Type(FixedSizeTypeMixin, TypeInfo):
cpp_type = "int32_t"
default_value = "0"
decode_32bit = "value.as_sfixed32()"
decode_expr = "value.as_sfixed32()"
encode_func = "encode_sfixed32"
wire_type = WireType.FIXED32 # Uses wire type 5
@@ -1615,7 +1585,6 @@ class SFixed32Type(FixedSizeTypeMixin, TypeInfo):
class SFixed64Type(FixedSizeTypeMixin, TypeInfo):
cpp_type = "int64_t"
default_value = "0"
decode_64bit = "value.as_sfixed64()"
encode_func = "encode_sfixed64"
wire_type = WireType.FIXED64 # Uses wire type 1
@@ -1642,7 +1611,7 @@ class SInt32Type(VarintTypeMixin, TypeInfo):
cpp_type = "int32_t"
_varint_max_bits = 32 # zigzag encoding keeps it 32-bit
default_value = "0"
decode_varint = "decode_zigzag32(static_cast<uint32_t>(value))"
decode_expr = "decode_zigzag32(static_cast<uint32_t>(value.as_varint()))"
encode_func = "encode_sint32"
wire_type = WireType.VARINT # Uses wire type 0
@@ -1663,7 +1632,7 @@ class SInt64Type(VarintTypeMixin, TypeInfo):
cpp_type = "int64_t"
_varint_max_bits = 64
default_value = "0"
decode_varint = "decode_zigzag64(value)"
decode_expr = "decode_zigzag64(value.as_varint())"
encode_func = "encode_sint64"
wire_type = WireType.VARINT # Uses wire type 0
@@ -1816,23 +1785,11 @@ 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,
if isinstance(self._ti, MessageType) and _is_inline_encode(self._ti.cpp_type):
return _generate_inline_encode_block(
self.number, self._ti.cpp_type, element
)
# 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 self._ti.encode_element(self.number, element)
@property
def cpp_type(self) -> str:
@@ -2126,55 +2083,23 @@ class RepeatedTypeInfo(TypeInfo):
return self._ti.wire_type
@property
def decode_varint_content(self) -> str:
# Pointer fields don't support decoding
if self._use_pointer:
return None
content = self._ti.decode_varint
if content is None:
return None
return (
f"case {self.number}: this->{self.field_name}.push_back({content}); break;"
)
def decode_expr(self) -> str | None:
return self._ti.decode_expr
def _decode_store(self, expr: str) -> str:
return f"this->{self.field_name}.push_back({expr});"
@property
def decode_length_content(self) -> str:
def decode_content(self) -> str | None:
# Pointer fields don't support decoding
if self._use_pointer:
return None
content = self._ti.decode_length
if content is None and isinstance(self._ti, MessageType):
# Special handling for non-template message decoding
return f"case {self.number}: this->{self.field_name}.emplace_back(); value.decode_to_message(this->{self.field_name}.back()); break;"
if content is None:
return None
return (
f"case {self.number}: this->{self.field_name}.push_back({content}); break;"
)
@property
def decode_32bit_content(self) -> str:
# Pointer fields don't support decoding
if self._use_pointer:
return None
content = self._ti.decode_32bit
if content is None:
return None
return (
f"case {self.number}: this->{self.field_name}.push_back({content}); break;"
)
@property
def decode_64bit_content(self) -> str:
# Pointer fields don't support decoding
if self._use_pointer:
return None
content = self._ti.decode_64bit
if content is None:
return None
return (
f"case {self.number}: this->{self.field_name}.push_back({content}); break;"
)
if isinstance(self._ti, MessageType):
return self.decode_case(
f"this->{self.field_name}.emplace_back();\n"
f"value.decode_to_message(this->{self.field_name}.back());"
)
return super().decode_content
@property
def _ti_is_bool(self) -> bool:
@@ -2182,20 +2107,7 @@ class RepeatedTypeInfo(TypeInfo):
return isinstance(self._ti, BoolType)
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,
)
# 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 self._ti.encode_element(self.number, element)
@property
def encode_content(self) -> str:
@@ -2590,10 +2502,7 @@ def build_message_type(
) -> tuple[str, str, str]:
public_content: list[str] = []
protected_content: list[str] = []
decode_varint: list[str] = []
decode_length: list[str] = []
decode_32bit: list[str] = []
decode_64bit: list[str] = []
decode: list[str] = []
encode: list[str] = []
dump: list[str] = []
size_calc: list[str] = []
@@ -2722,22 +2631,8 @@ def build_message_type(
if field.options.HasExtension(pb.field_ifdef):
field_ifdef = field.options.Extensions[pb.field_ifdef]
if ti.decode_varint_content:
decode_varint.extend(
wrap_with_ifdef(ti.decode_varint_content, field_ifdef)
)
if ti.decode_length_content:
decode_length.extend(
wrap_with_ifdef(ti.decode_length_content, field_ifdef)
)
if ti.decode_32bit_content:
decode_32bit.extend(
wrap_with_ifdef(ti.decode_32bit_content, field_ifdef)
)
if ti.decode_64bit_content:
decode_64bit.extend(
wrap_with_ifdef(ti.decode_64bit_content, field_ifdef)
)
if case := ti.decode_content:
decode.extend(wrap_with_ifdef(case, field_ifdef))
if ti.dump_content:
# Check for field_ifdef option for dump as well
field_ifdef = None
@@ -2747,49 +2642,15 @@ def build_message_type(
dump.extend(wrap_with_ifdef(ti.dump_content, field_ifdef))
cpp = ""
if decode_varint:
o = f"bool {desc.name}::decode_varint(uint32_t field_id, proto_varint_value_t value) {{\n"
o += " switch (field_id) {\n"
o += indent("\n".join(decode_varint), " ") + "\n"
o += " default: return false;\n"
if decode:
o = f"void {desc.name}::decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) {{\n"
o += " const ProtoFieldValue value(data, scalar);\n"
o += " switch (tag) {\n"
o += indent("\n".join(decode), " ") + "\n"
o += " }\n"
o += " return true;\n"
o += "}\n"
cpp += o
prot = "bool decode_varint(uint32_t field_id, proto_varint_value_t value) override;"
protected_content.insert(0, prot)
if decode_length:
o = f"bool {desc.name}::decode_length(uint32_t field_id, ProtoLengthDelimited value) {{\n"
o += " switch (field_id) {\n"
o += indent("\n".join(decode_length), " ") + "\n"
o += " default: return false;\n"
o += " }\n"
o += " return true;\n"
o += "}\n"
cpp += o
prot = "bool decode_length(uint32_t field_id, ProtoLengthDelimited value) override;"
protected_content.insert(0, prot)
if decode_32bit:
o = f"bool {desc.name}::decode_32bit(uint32_t field_id, Proto32Bit value) {{\n"
o += " switch (field_id) {\n"
o += indent("\n".join(decode_32bit), " ") + "\n"
o += " default: return false;\n"
o += " }\n"
o += " return true;\n"
o += "}\n"
cpp += o
prot = "bool decode_32bit(uint32_t field_id, Proto32Bit value) override;"
protected_content.insert(0, prot)
if decode_64bit:
o = f"bool {desc.name}::decode_64bit(uint32_t field_id, Proto64Bit value) {{\n"
o += " switch (field_id) {\n"
o += indent("\n".join(decode_64bit), " ") + "\n"
o += " default: return false;\n"
o += " }\n"
o += " return true;\n"
o += "}\n"
cpp += o
prot = "bool decode_64bit(uint32_t field_id, Proto64Bit value) override;"
prot = "void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;"
protected_content.insert(0, prot)
# Generate custom decode() override for messages with FixedVector fields
@@ -249,7 +249,7 @@ static APIBuffer build_infrared_rf_transmit_wire() {
std::memcpy(bytes + len, packed, packed_len);
len += packed_len;
// field 6: modulation = 1 (non-zero so it's actually emitted and exercises
// decode_varint for this field, matching the documented layout above).
// decode_field for this field, matching the documented layout above).
put_byte(0x30);
put_varint(1);
@@ -0,0 +1,43 @@
esphome:
name: api-decode-wire-types-test
host:
api:
logger:
level: DEBUG
switch:
- platform: template
name: "Wire Switch"
optimistic: true
output:
- platform: template
id: wire_dim
type: float
write_action:
- lambda: ""
light:
- platform: monochromatic
name: "Wire Light"
output: wire_dim
default_transition_length: 0s
effects:
- pulse:
name: Pulse
text:
- platform: template
name: "Wire Text"
optimistic: true
mode: text
min_length: 0
max_length: 255
number:
- platform: template
name: "Wire Number"
optimistic: true
min_value: -1000
max_value: 1000
step: 0.5
+5 -4
View File
@@ -125,11 +125,12 @@ class RawApiClient:
await self.read_until_frame(MESSAGE_TYPE_OF[api_pb2.HelloResponse])
async def send_message(self, msg: message.Message) -> None:
await self.send_raw(MESSAGE_TYPE_OF[type(msg)], msg.SerializeToString())
async def send_raw(self, msg_type: int, payload: bytes) -> None:
"""Send a frame with a hand built payload, for shapes protobuf will not serialize."""
loop = asyncio.get_running_loop()
await loop.sock_sendall(
self._sock,
encode_frame(MESSAGE_TYPE_OF[type(msg)], msg.SerializeToString()),
)
await loop.sock_sendall(self._sock, encode_frame(msg_type, payload))
async def read_until_frame(self, msg_type: int, timeout: float = 10.0) -> None:
"""Read until at least one frame of msg_type has been received."""
@@ -0,0 +1,142 @@
"""decode_field() must take fields that match their declared wire type, drop the ones that do
not, skip unknown fields, and handle two byte tags, varints and length prefixes."""
from __future__ import annotations
from collections.abc import Callable
import struct
from aioesphomeapi import (
EntityState,
LightState,
NumberState,
SwitchState,
TextState,
api_pb2,
)
import pytest
from .raw_api_client import MESSAGE_TYPE_OF, RawApiClient, encode_varint
from .state_utils import InitialStateHelper, StateWaiter, require_entity
from .types import APIClientConnectedFactory, RunCompiledFunction
SWITCH_COMMAND = MESSAGE_TYPE_OF[api_pb2.SwitchCommandRequest]
WIRE_VARINT, WIRE_LENGTH, WIRE_FIXED32 = 0, 2, 5
def tag(field: int, wire_type: int) -> bytes:
return encode_varint((field << 3) | wire_type)
@pytest.mark.asyncio
async def test_api_decode_wire_types(
yaml_config: str,
run_compiled: RunCompiledFunction,
api_client_connected: APIClientConnectedFactory,
unused_tcp_port: int,
) -> None:
async with (
run_compiled(yaml_config),
api_client_connected() as client,
RawApiClient(unused_tcp_port) as raw,
):
entities, _ = await client.list_entities_services()
switch = require_entity(entities, "wire_switch")
light = require_entity(entities, "wire_light")
text = require_entity(entities, "wire_text")
number = require_entity(entities, "wire_number")
key = tag(1, WIRE_FIXED32) + struct.pack("<I", switch.key)
on, off = tag(2, WIRE_VARINT) + b"\x01", tag(2, WIRE_VARINT) + b"\x00"
switch_states: list[bool] = []
waiter = StateWaiter()
def on_state(state: EntityState) -> None:
if isinstance(state, SwitchState) and state.key == switch.key:
switch_states.append(state.state)
waiter.on_state(state)
def switch_is(value: bool) -> Callable[[EntityState], bool]:
return lambda s: (
isinstance(s, SwitchState) and s.key == switch.key and s.state is value
)
def number_is(value: float) -> Callable[[EntityState], bool]:
return lambda s: (
isinstance(s, NumberState) and s.key == number.key and s.state == value
)
initial = InitialStateHelper(entities)
client.subscribe_states(initial.on_state_wrapper(on_state))
await initial.wait_for_initial_states()
await raw.connect()
# A well formed command: fixed32 key, varint state
await raw.send_raw(SWITCH_COMMAND, key + on)
await waiter.expect(switch_is(True))
await raw.send_raw(SWITCH_COMMAND, key + off)
await waiter.expect(switch_is(False))
# The same field with the wrong wire type is dropped, and a varint key never matches an
# entity; each of these would turn the switch on if the payload were read as a varint
seen = len(switch_states)
await raw.send_raw(SWITCH_COMMAND, key + tag(2, WIRE_LENGTH) + b"\x01\x01")
await raw.send_raw(
SWITCH_COMMAND, key + tag(2, WIRE_FIXED32) + b"\x01\x00\x00\x00"
)
await raw.send_raw(
SWITCH_COMMAND, tag(1, WIRE_VARINT) + encode_varint(switch.key) + on
)
# Ordered on the raw socket itself: this frame cannot be parsed before the bad ones, so
# the only switch state since the marker must be the one it produces
await raw.send_raw(SWITCH_COMMAND, key + on)
await waiter.expect(switch_is(True), label="switch on after wrong wire types")
assert switch_states[seen:] == [True]
await raw.send_raw(SWITCH_COMMAND, key + off)
await waiter.expect(switch_is(False))
# Truncated bodies stop the decode loop without taking the connection down: a tag with its
# continuation bit set and nothing after it, a length prefix past the end of the payload,
# and a fixed32 with two of its four bytes
seen = len(switch_states)
await raw.send_raw(SWITCH_COMMAND, key + b"\x80")
await raw.send_raw(SWITCH_COMMAND, key + tag(2, WIRE_LENGTH) + b"\x7f" + b"ab")
await raw.send_raw(SWITCH_COMMAND, tag(1, WIRE_FIXED32) + b"\x01\x02")
await raw.send_raw(SWITCH_COMMAND, key + on)
await waiter.expect(switch_is(True), label="switch on after truncated frames")
assert switch_states[seen:] == [True]
await raw.send_raw(SWITCH_COMMAND, key + off)
await waiter.expect(switch_is(False))
# A negative number goes through the fixed32 float path of a normal client
client.number_command(number.key, -77.5)
await waiter.expect(number_is(-77.5))
# An unknown field ahead of the known ones is skipped; field 200 needs a two byte tag
await raw.send_raw(
SWITCH_COMMAND, tag(200, WIRE_VARINT) + encode_varint(300) + key + on
)
await waiter.expect(switch_is(True))
# Two byte tags (effect fields 18 and 19) and a two byte varint (300 ms transition)
client.light_command(
light.key, state=True, brightness=0.5, transition_length=0.3, effect="Pulse"
)
await waiter.expect(
lambda s: (
isinstance(s, LightState) and s.key == light.key and s.effect == "Pulse"
)
)
client.light_command(light.key, effect="None", state=False)
await waiter.expect(
lambda s: isinstance(s, LightState) and s.key == light.key and not s.state
)
# A string whose length prefix needs two varint bytes
long_text = "w" * 200
client.text_command(text.key, long_text)
await waiter.expect(
lambda s: (
isinstance(s, TextState) and s.key == text.key and s.state == long_text
)
)
@@ -18,7 +18,9 @@ sys.path.insert(0, str(Path(__file__).parents[4] / "script" / "api_protobuf"))
import aioesphomeapi.api_options_pb2 as pb # noqa: E402
from api_protobuf import ( # noqa: E402
MAX_MESSAGE_ID,
SOURCE_CLIENT,
_make_ifdef_line,
build_message_type,
create_field_type_info,
get_varint64_ifdef,
validate_message_id,
@@ -36,12 +38,15 @@ def _file_with_messages(
file_desc = descriptor_pb2.FileDescriptorProto(name="test.proto")
for name, field_type, deprecated in messages:
msg = file_desc.message_type.add(name=name)
field = msg.field.add(name="value", number=1, type=field_type)
field = msg.field.add()
field.CopyFrom(_field(field_type))
field.options.deprecated = deprecated
return file_desc
UINT64 = descriptor_pb2.FieldDescriptorProto.TYPE_UINT64
MESSAGE = descriptor_pb2.FieldDescriptorProto.TYPE_MESSAGE
DOUBLE = descriptor_pb2.FieldDescriptorProto.TYPE_DOUBLE
INT64 = descriptor_pb2.FieldDescriptorProto.TYPE_INT64
SINT64 = descriptor_pb2.FieldDescriptorProto.TYPE_SINT64
UINT32 = descriptor_pb2.FieldDescriptorProto.TYPE_UINT32
@@ -182,3 +187,106 @@ def test_multi_byte_tag_fixed32_falls_back_to_the_generic_helper(
content = _encode_field(field_type, number=16)
assert "write_tag_and_fixed32" not in content, content
assert content.startswith("pos = ProtoEncode::encode_"), content
def _decode_case(field_type: int, number: int, *, repeated: bool = False) -> str:
"""Return the decode_field() case the generator emits for one decoded field."""
field = _field(field_type, number, repeated=repeated)
if field_type == MESSAGE:
field.type_name = ".Sub"
return create_field_type_info(
field, needs_decode=True, needs_encode=False
).decode_content
@pytest.mark.parametrize(
("field_type", "number", "wire_type", "accessor"),
[
(UINT32, 2, "WIRE_TYPE_VARINT", "value.as_varint()"),
(BOOL, 3, "WIRE_TYPE_VARINT", "value.as_bool()"),
(STRING, 1, "WIRE_TYPE_LENGTH_DELIMITED", "value.data()"),
(FLOAT, 4, "WIRE_TYPE_FIXED32", "value.as_float()"),
(FIXED32, 5, "WIRE_TYPE_FIXED32", "value.as_fixed32()"),
],
)
def test_decode_cases_carry_field_number_and_wire_type(
field_type: int, number: int, wire_type: str, accessor: str
) -> None:
"""Each decoded field yields one case keyed on its number and declared wire type."""
case = _decode_case(field_type, number)
lines = case.splitlines()
assert lines[0] == f"case proto_tag({number}, {wire_type}):", case
assert accessor in lines[1], case
assert lines[-1].strip() == "break;", case
@pytest.mark.parametrize(
("field_type", "repeated", "wire_type", "store"),
[
(UINT32, True, "WIRE_TYPE_VARINT", "this->value.push_back(value.as_varint());"),
(
STRING,
True,
"WIRE_TYPE_LENGTH_DELIMITED",
"this->value.push_back(value.as_string());",
),
(
MESSAGE,
False,
"WIRE_TYPE_LENGTH_DELIMITED",
"value.decode_to_message(this->value);",
),
(
MESSAGE,
True,
"WIRE_TYPE_LENGTH_DELIMITED",
"value.decode_to_message(this->value.back());",
),
],
)
def test_repeated_and_message_fields_decode_through_the_same_case_shape(
field_type: int, repeated: bool, wire_type: str, store: str
) -> None:
"""Repeated and sub message fields land in the one switch with their own store."""
case = _decode_case(field_type, 7, repeated=repeated)
lines = case.splitlines()
assert lines[0] == f"case proto_tag(7, {wire_type}):", case
assert store in case, case
if field_type == MESSAGE and repeated:
assert "this->value.emplace_back();" in case, case
assert lines[-1].strip() == "break;", case
def test_a_fixed64_field_fails_at_generation_time() -> None:
"""The decode loop has no 64 bit wire type path, so such a field must never reach it silently."""
desc = descriptor_pb2.DescriptorProto(name="Wide")
desc.field.add(name="ratio", number=1, type=DOUBLE)
with pytest.raises(
ValueError, match="64-bit type 'double' .*ratio.* not supported"
):
build_message_type(desc, {}, {"Wide": SOURCE_CLIENT})
def test_message_gets_a_single_decode_field_override() -> None:
"""All wire types of a decoded message land in one decode_field() switch."""
desc = descriptor_pb2.DescriptorProto(name="Mixed")
desc.field.add(name="name", number=1, type=STRING)
desc.field.add(name="count", number=2, type=UINT32)
desc.field.add(name="level", number=3, type=FLOAT)
header, cpp, _ = build_message_type(desc, {}, {"Mixed": SOURCE_CLIENT})
decl = "void decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;"
assert header.count(decl) == 1
assert (
cpp.count(
"void Mixed::decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) {"
)
== 1
)
assert "switch (tag) {" in cpp
assert "const ProtoFieldValue value(data, scalar);" in cpp
for number, wire_type in (
(1, "WIRE_TYPE_LENGTH_DELIMITED"),
(2, "WIRE_TYPE_VARINT"),
(3, "WIRE_TYPE_FIXED32"),
):
assert f"case proto_tag({number}, {wire_type}):" in cpp, cpp