mirror of
https://github.com/esphome/esphome.git
synced 2026-09-11 15:27:33 +00:00
[api] Collapse the three protobuf decode virtuals into one
Every decodable message overrode up to three virtuals, one per wire type, so each carried a five slot vtable and up to three functions with their own prologue and return tails. The shared decode loop now parses the payload for the wire type into a ProtoFieldValue and calls a single decode_field() virtual with the tag, the field number and the wire type; the generated override is one switch. The switch key is chosen per target through PROTO_DECODE_KEY. Embedded builds compile switches to compare chains (ESP-IDF passes -fno-jump-tables), so they key on the full wire tag, one compare per field with no separate wire type check. The host compiler builds a jump table for the dense field number switch, so there the key is the field number and PROTO_DECODE_GUARD rejects a mismatched wire type. Both forms drop a field that arrives with a wire type it does not declare, exactly as the per wire type virtuals did. Per decodable message the vtable shrinks from 20 to 12 bytes on xtensa and the extra decode functions fold into one; the shared loop shrinks as well. Host instruction counts per decoded field are unchanged apart from the guard compare, which replaces the prologue of the separate function it used to call.
This commit is contained in:
@@ -229,12 +229,27 @@ class TypeInfo(ABC):
|
||||
def class_member(self) -> str:
|
||||
return f"{self.cpp_type} {self.field_name}{{{self.default_value}}};"
|
||||
|
||||
# decode_field() cases are keyed through the PROTO_DECODE_* macros in proto.h: embedded
|
||||
# targets switch on the full wire tag, the host switches on the field number and guards
|
||||
# the wire type. Either way a field that arrives with the wrong wire type falls through
|
||||
# to "return false" instead of being read from the wrong ProtoFieldValue member.
|
||||
def decode_case(self, wire_type: WireType, body: str) -> str:
|
||||
"""Emit one decode_field() case for a field and the wire type it expects."""
|
||||
return (
|
||||
f"case PROTO_DECODE_CASE({self.number}, {int(wire_type)}):\n"
|
||||
f" PROTO_DECODE_GUARD(wire_type, {int(wire_type)});\n"
|
||||
f" {body}\n"
|
||||
f" break;"
|
||||
)
|
||||
|
||||
@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;"
|
||||
return self.decode_case(
|
||||
WireType.VARINT, f"this->{self.field_name} = {content};"
|
||||
)
|
||||
|
||||
decode_varint = None
|
||||
|
||||
@@ -243,7 +258,9 @@ class TypeInfo(ABC):
|
||||
content = self.decode_length
|
||||
if content is None:
|
||||
return None
|
||||
return f"case {self.number}: this->{self.field_name} = {content}; break;"
|
||||
return self.decode_case(
|
||||
WireType.LENGTH_DELIMITED, f"this->{self.field_name} = {content};"
|
||||
)
|
||||
|
||||
decode_length = None
|
||||
|
||||
@@ -252,19 +269,12 @@ class TypeInfo(ABC):
|
||||
content = self.decode_32bit
|
||||
if content is None:
|
||||
return None
|
||||
return f"case {self.number}: this->{self.field_name} = {content}; break;"
|
||||
return self.decode_case(
|
||||
WireType.FIXED32, f"this->{self.field_name} = {content};"
|
||||
)
|
||||
|
||||
decode_32bit = None
|
||||
|
||||
@property
|
||||
def decode_64bit_content(self) -> str:
|
||||
content = self.decode_64bit
|
||||
if content is None:
|
||||
return None
|
||||
return f"case {self.number}: this->{self.field_name} = {content}; break;"
|
||||
|
||||
decode_64bit = None
|
||||
|
||||
# Mapping from encode_func to raw encode expression template.
|
||||
# When a forced field has a single-byte tag, the code generator emits
|
||||
# write_raw_byte(tag) + raw encode instead of the full encode_* method,
|
||||
@@ -638,7 +648,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
|
||||
|
||||
@@ -693,7 +702,7 @@ class Int64Type(VarintTypeMixin, TypeInfo):
|
||||
cpp_type = "int64_t"
|
||||
_varint_max_bits = 64
|
||||
default_value = "0"
|
||||
decode_varint = "static_cast<int64_t>(value)"
|
||||
decode_varint = "static_cast<int64_t>(value.as_varint())"
|
||||
encode_func = "encode_int64"
|
||||
wire_type = WireType.VARINT # Uses wire type 0
|
||||
|
||||
@@ -714,7 +723,7 @@ class UInt64Type(VarintTypeMixin, TypeInfo):
|
||||
cpp_type = "uint64_t"
|
||||
_varint_max_bits = 64
|
||||
default_value = "0"
|
||||
decode_varint = "value"
|
||||
decode_varint = "value.as_varint()"
|
||||
encode_func = "encode_uint64"
|
||||
wire_type = WireType.VARINT # Uses wire type 0
|
||||
|
||||
@@ -749,7 +758,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_varint = "static_cast<int32_t>(value.as_varint())"
|
||||
encode_func = "encode_int32"
|
||||
wire_type = WireType.VARINT # Uses wire type 0
|
||||
|
||||
@@ -769,7 +778,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
|
||||
|
||||
@@ -824,7 +832,7 @@ class BoolType(VarintTypeMixin, TypeInfo):
|
||||
_varint_max_bits = 1
|
||||
cpp_type = "bool"
|
||||
default_value = "false"
|
||||
decode_varint = "value != 0"
|
||||
decode_varint = "value.as_varint() != 0"
|
||||
encode_func = "encode_bool"
|
||||
wire_type = WireType.VARINT # Uses wire type 0
|
||||
|
||||
@@ -1014,13 +1022,15 @@ class MessageType(TypeInfo):
|
||||
# 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 self.decode_case(
|
||||
WireType.LENGTH_DELIMITED,
|
||||
f"value.decode_to_message(this->{self.field_name});\n"
|
||||
f" this->has_{self.name} = true;",
|
||||
)
|
||||
return f"case {self.number}: value.decode_to_message(this->{self.field_name}); break;"
|
||||
return self.decode_case(
|
||||
WireType.LENGTH_DELIMITED,
|
||||
f"value.decode_to_message(this->{self.field_name});",
|
||||
)
|
||||
|
||||
def dump(self, name: str) -> str:
|
||||
return f"{name}.dump_to(out);"
|
||||
@@ -1216,11 +1226,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;
|
||||
}}"""
|
||||
return self.decode_case(
|
||||
WireType.LENGTH_DELIMITED,
|
||||
f"this->{self.field_name} = value.data();\n"
|
||||
f" this->{self.field_name}_len = value.size();",
|
||||
)
|
||||
|
||||
def dump(self, name: str) -> str:
|
||||
return (
|
||||
@@ -1281,10 +1291,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;
|
||||
}}"""
|
||||
return self.decode_case(
|
||||
WireType.LENGTH_DELIMITED,
|
||||
f"this->{self.field_name} = StringRef(reinterpret_cast<const char *>(value.data()), value.size());",
|
||||
)
|
||||
|
||||
def dump(self, name: str) -> str:
|
||||
# Not used since we use dump_field, but required by abstract base class
|
||||
@@ -1355,12 +1365,12 @@ class PackedBufferTypeInfo(TypeInfo):
|
||||
@property
|
||||
def decode_length_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(
|
||||
WireType.LENGTH_DELIMITED,
|
||||
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:
|
||||
@@ -1446,16 +1456,13 @@ 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
|
||||
body = "const std::string &data_str = value.as_string();\n"
|
||||
body += f" this->{self.field_name}_len = data_str.size();\n"
|
||||
body += f" if (this->{self.field_name}_len > {self.array_size}) {{\n"
|
||||
body += f" this->{self.field_name}_len = {self.array_size};\n"
|
||||
body += " }\n"
|
||||
body += f" memcpy(this->{self.field_name}, data_str.data(), this->{self.field_name}_len);"
|
||||
return self.decode_case(WireType.LENGTH_DELIMITED, body)
|
||||
|
||||
@property
|
||||
def encode_content(self) -> str:
|
||||
@@ -1518,7 +1525,7 @@ class UInt32Type(VarintTypeMixin, TypeInfo):
|
||||
cpp_type = "uint32_t"
|
||||
_varint_max_bits = 32
|
||||
default_value = "0"
|
||||
decode_varint = "value"
|
||||
decode_varint = "value.as_varint()"
|
||||
encode_func = "encode_uint32"
|
||||
wire_type = WireType.VARINT # Uses wire type 0
|
||||
|
||||
@@ -1547,7 +1554,7 @@ class EnumType(VarintTypeMixin, TypeInfo):
|
||||
|
||||
@property
|
||||
def decode_varint(self) -> str:
|
||||
return f"static_cast<{self.cpp_type}>(value)"
|
||||
return f"static_cast<{self.cpp_type}>(value.as_varint())"
|
||||
|
||||
default_value = ""
|
||||
wire_type = WireType.VARINT # Uses wire type 0
|
||||
@@ -1620,7 +1627,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
|
||||
|
||||
@@ -1647,7 +1653,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_varint = "decode_zigzag32(static_cast<uint32_t>(value.as_varint()))"
|
||||
encode_func = "encode_sint32"
|
||||
wire_type = WireType.VARINT # Uses wire type 0
|
||||
|
||||
@@ -1668,7 +1674,7 @@ class SInt64Type(VarintTypeMixin, TypeInfo):
|
||||
cpp_type = "int64_t"
|
||||
_varint_max_bits = 64
|
||||
default_value = "0"
|
||||
decode_varint = "decode_zigzag64(value)"
|
||||
decode_varint = "decode_zigzag64(value.as_varint())"
|
||||
encode_func = "encode_sint64"
|
||||
wire_type = WireType.VARINT # Uses wire type 0
|
||||
|
||||
@@ -2138,8 +2144,8 @@ class RepeatedTypeInfo(TypeInfo):
|
||||
content = self._ti.decode_varint
|
||||
if content is None:
|
||||
return None
|
||||
return (
|
||||
f"case {self.number}: this->{self.field_name}.push_back({content}); break;"
|
||||
return self.decode_case(
|
||||
WireType.VARINT, f"this->{self.field_name}.push_back({content});"
|
||||
)
|
||||
|
||||
@property
|
||||
@@ -2150,11 +2156,15 @@ class RepeatedTypeInfo(TypeInfo):
|
||||
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;"
|
||||
return self.decode_case(
|
||||
WireType.LENGTH_DELIMITED,
|
||||
f"this->{self.field_name}.emplace_back();\n"
|
||||
f" value.decode_to_message(this->{self.field_name}.back());",
|
||||
)
|
||||
if content is None:
|
||||
return None
|
||||
return (
|
||||
f"case {self.number}: this->{self.field_name}.push_back({content}); break;"
|
||||
return self.decode_case(
|
||||
WireType.LENGTH_DELIMITED, f"this->{self.field_name}.push_back({content});"
|
||||
)
|
||||
|
||||
@property
|
||||
@@ -2165,20 +2175,8 @@ class RepeatedTypeInfo(TypeInfo):
|
||||
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;"
|
||||
return self.decode_case(
|
||||
WireType.FIXED32, f"this->{self.field_name}.push_back({content});"
|
||||
)
|
||||
|
||||
@property
|
||||
@@ -2595,10 +2593,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] = []
|
||||
@@ -2727,22 +2722,13 @@ 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)
|
||||
)
|
||||
for case in (
|
||||
ti.decode_varint_content,
|
||||
ti.decode_length_content,
|
||||
ti.decode_32bit_content,
|
||||
):
|
||||
if case:
|
||||
decode.extend(wrap_with_ifdef(case, field_ifdef))
|
||||
if ti.dump_content:
|
||||
# Check for field_ifdef option for dump as well
|
||||
field_ifdef = None
|
||||
@@ -2752,49 +2738,18 @@ 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"
|
||||
if decode:
|
||||
# One virtual per message: the shared decode loop parses the payload for the wire
|
||||
# type and hands it over with the tag, so a single switch covers every field.
|
||||
o = f"bool {desc.name}::decode_field(uint32_t tag, uint32_t field_id, uint32_t wire_type, ProtoFieldValue value) {{\n"
|
||||
o += " switch (PROTO_DECODE_KEY(tag, field_id)) {\n"
|
||||
o += indent("\n".join(decode), " ") + "\n"
|
||||
o += " default: return false;\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 = "bool decode_field(uint32_t tag, uint32_t field_id, uint32_t wire_type, ProtoFieldValue value) override;"
|
||||
protected_content.insert(0, prot)
|
||||
|
||||
# Generate custom decode() override for messages with FixedVector fields
|
||||
|
||||
Reference in New Issue
Block a user