mirror of
https://github.com/esphome/esphome.git
synced 2026-09-11 15:27:33 +00:00
[api] Emit one decode case per field from a single generator property
With one decode_field() switch per message, the three per wire type content properties only differed in the attribute they read; a single decode_content built from decode_expr() replaces them, and repeated fields reuse the element type's expression. Case bodies with several statements get their block from the body itself instead of a caller flag, the fixed byte array body copies straight from the payload instead of through a heap std::string, and the decode comments no longer restate the switch keying explained next to the macros.
This commit is contained in:
@@ -748,10 +748,8 @@ class ProtoDecodableMessage : public ProtoMessage {
|
||||
~ProtoDecodableMessage() = default;
|
||||
/// Store one decoded field. \p tag is the wire tag (field number and wire type), \p data points at
|
||||
/// the field payload and \p scalar is the varint or fixed32 value, or the payload length for a
|
||||
/// length-delimited field. Three register arguments keep the shared loop free of spills; the
|
||||
/// generated override wraps them in a ProtoFieldValue and keys its switch through
|
||||
/// PROTO_DECODE_KEY. Overrides reject a field that arrived with a wire type other than the one it
|
||||
/// declares. Return false for unknown or mismatched fields.
|
||||
/// length-delimited field. Three register arguments keep the shared loop free of spills. Return
|
||||
/// false for an unknown field or one that arrived with a wire type it does not declare.
|
||||
/// One virtual instead of one per wire type keeps each message's vtable at a single slot.
|
||||
// NOTE: wire type 1 (64-bit fixed) is not supported
|
||||
virtual bool decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) { return false; }
|
||||
|
||||
@@ -229,54 +229,45 @@ 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, scoped: bool = False) -> str:
|
||||
# Cases are keyed through the PROTO_DECODE_* macros in proto.h, which is where the
|
||||
# host and embedded switch shapes are explained.
|
||||
def decode_case(self, wire_type: WireType, body: str) -> str:
|
||||
"""Emit one decode_field() case for a field and the wire type it expects.
|
||||
|
||||
Bodies that declare locals must be scoped so the jump to the next case label
|
||||
does not cross an initialization.
|
||||
Multi-statement bodies get their own block so a local in one case cannot be
|
||||
jumped over by a later case label.
|
||||
"""
|
||||
label = f"case PROTO_DECODE_CASE({self.number}, {int(wire_type)}):"
|
||||
guard = f"PROTO_DECODE_GUARD(tag, {self.number}, {int(wire_type)});"
|
||||
if scoped:
|
||||
if "\n" in body:
|
||||
return f"{label} {{\n" + indent(f"{guard}\n{body}\nbreak;") + "\n}"
|
||||
return f"{label}\n" + indent(f"{guard}\n{body}\nbreak;")
|
||||
|
||||
@property
|
||||
def decode_varint_content(self) -> str:
|
||||
content = self.decode_varint
|
||||
if content is None:
|
||||
return None
|
||||
return self.decode_case(
|
||||
WireType.VARINT, f"this->{self.field_name} = {content};"
|
||||
)
|
||||
|
||||
# Value expression that decodes this type from the ProtoFieldValue, per wire type.
|
||||
# A type sets exactly one of them; None everywhere means the field is never decoded.
|
||||
decode_varint = None
|
||||
|
||||
@property
|
||||
def decode_length_content(self) -> str:
|
||||
content = self.decode_length
|
||||
if content is None:
|
||||
return None
|
||||
return self.decode_case(
|
||||
WireType.LENGTH_DELIMITED, f"this->{self.field_name} = {content};"
|
||||
)
|
||||
|
||||
decode_length = None
|
||||
decode_32bit = None
|
||||
|
||||
def decode_expr(self) -> tuple[WireType, str] | None:
|
||||
"""Wire type and value expression for decoding this field, or None."""
|
||||
for wire_type, content in (
|
||||
(WireType.VARINT, self.decode_varint),
|
||||
(WireType.LENGTH_DELIMITED, self.decode_length),
|
||||
(WireType.FIXED32, self.decode_32bit),
|
||||
):
|
||||
if content is not None:
|
||||
return wire_type, content
|
||||
return None
|
||||
|
||||
@property
|
||||
def decode_32bit_content(self) -> str:
|
||||
content = self.decode_32bit
|
||||
if content is 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()
|
||||
if expr is None:
|
||||
return None
|
||||
return self.decode_case(
|
||||
WireType.FIXED32, f"this->{self.field_name} = {content};"
|
||||
)
|
||||
|
||||
decode_32bit = None
|
||||
wire_type, content = expr
|
||||
return self.decode_case(wire_type, f"this->{self.field_name} = {content};")
|
||||
|
||||
# Mapping from encode_func to raw encode expression template.
|
||||
# When a forced field has a single-byte tag, the code generator emits
|
||||
@@ -1019,7 +1010,7 @@ class MessageType(TypeInfo):
|
||||
)
|
||||
|
||||
@property
|
||||
def decode_length_content(self) -> str:
|
||||
def decode_content(self) -> str:
|
||||
# Custom decode that doesn't use templates
|
||||
if self._track_presence:
|
||||
# decode_to_message() cannot report failure, so setting the flag
|
||||
@@ -1029,7 +1020,6 @@ class MessageType(TypeInfo):
|
||||
WireType.LENGTH_DELIMITED,
|
||||
f"value.decode_to_message(this->{self.field_name});\n"
|
||||
f"this->has_{self.name} = true;",
|
||||
scoped=True,
|
||||
)
|
||||
return self.decode_case(
|
||||
WireType.LENGTH_DELIMITED,
|
||||
@@ -1185,7 +1175,7 @@ class PointerToBufferTypeBase(TypeInfo):
|
||||
|
||||
@property
|
||||
def decode_length(self) -> str | None:
|
||||
# This is handled in decode_length_content
|
||||
# This is handled in decode_content
|
||||
return None
|
||||
|
||||
@property
|
||||
@@ -1229,12 +1219,11 @@ class PointerToBytesBufferType(PointerToBufferTypeBase):
|
||||
)
|
||||
|
||||
@property
|
||||
def decode_length_content(self) -> str | None:
|
||||
def decode_content(self) -> str:
|
||||
return self.decode_case(
|
||||
WireType.LENGTH_DELIMITED,
|
||||
f"this->{self.field_name} = value.data();\n"
|
||||
f"this->{self.field_name}_len = value.size();",
|
||||
scoped=True,
|
||||
)
|
||||
|
||||
def dump(self, name: str) -> str:
|
||||
@@ -1295,7 +1284,7 @@ class PointerToStringBufferType(PointerToBufferTypeBase):
|
||||
)
|
||||
|
||||
@property
|
||||
def decode_length_content(self) -> str | None:
|
||||
def decode_content(self) -> str:
|
||||
return self.decode_case(
|
||||
WireType.LENGTH_DELIMITED,
|
||||
f"this->{self.field_name} = StringRef(reinterpret_cast<const char *>(value.data()), value.size());",
|
||||
@@ -1368,14 +1357,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 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());",
|
||||
scoped=True,
|
||||
)
|
||||
|
||||
@property
|
||||
@@ -1461,14 +1449,12 @@ class FixedArrayBytesType(TypeInfo):
|
||||
]
|
||||
|
||||
@property
|
||||
def decode_length_content(self) -> str:
|
||||
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, scoped=True)
|
||||
def decode_content(self) -> str:
|
||||
return self.decode_case(
|
||||
WireType.LENGTH_DELIMITED,
|
||||
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:
|
||||
@@ -2143,47 +2129,23 @@ class RepeatedTypeInfo(TypeInfo):
|
||||
return self._ti.wire_type
|
||||
|
||||
@property
|
||||
def decode_varint_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_varint
|
||||
if content is None:
|
||||
return None
|
||||
return self.decode_case(
|
||||
WireType.VARINT, f"this->{self.field_name}.push_back({content});"
|
||||
)
|
||||
|
||||
@property
|
||||
def decode_length_content(self) -> str:
|
||||
# Pointer fields don't support decoding
|
||||
if self._use_pointer:
|
||||
return None
|
||||
content = self._ti.decode_length
|
||||
if content is None and isinstance(self._ti, MessageType):
|
||||
if isinstance(self._ti, MessageType):
|
||||
# Special handling for non-template message decoding
|
||||
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());",
|
||||
scoped=True,
|
||||
)
|
||||
if content is None:
|
||||
expr = self._ti.decode_expr()
|
||||
if expr is None:
|
||||
return None
|
||||
wire_type, content = expr
|
||||
return self.decode_case(
|
||||
WireType.LENGTH_DELIMITED, f"this->{self.field_name}.push_back({content});"
|
||||
)
|
||||
|
||||
@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 self.decode_case(
|
||||
WireType.FIXED32, f"this->{self.field_name}.push_back({content});"
|
||||
wire_type, f"this->{self.field_name}.push_back({content});"
|
||||
)
|
||||
|
||||
@property
|
||||
@@ -2729,13 +2691,8 @@ def build_message_type(
|
||||
if field.options.HasExtension(pb.field_ifdef):
|
||||
field_ifdef = field.options.Extensions[pb.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 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
|
||||
|
||||
@@ -186,34 +186,21 @@ def test_multi_byte_tag_fixed32_falls_back_to_the_generic_helper(
|
||||
assert content.startswith("pos = ProtoEncode::encode_"), content
|
||||
|
||||
|
||||
def _decode_cases(field_type: int, number: int) -> list[str]:
|
||||
"""Return the decode_field() case lines the generator emits for one decoded field."""
|
||||
def _decode_case(field_type: int, number: int) -> str:
|
||||
"""Return the decode_field() case the generator emits for one decoded field."""
|
||||
field = descriptor_pb2.FieldDescriptorProto(
|
||||
name="value", number=number, type=field_type
|
||||
)
|
||||
ti = create_field_type_info(field, needs_decode=True, needs_encode=False)
|
||||
return [
|
||||
case
|
||||
for case in (
|
||||
ti.decode_varint_content,
|
||||
ti.decode_length_content,
|
||||
ti.decode_32bit_content,
|
||||
)
|
||||
if case
|
||||
]
|
||||
|
||||
|
||||
UINT32_T = descriptor_pb2.FieldDescriptorProto.TYPE_UINT32
|
||||
STRING_T = descriptor_pb2.FieldDescriptorProto.TYPE_STRING
|
||||
BOOL_T = descriptor_pb2.FieldDescriptorProto.TYPE_BOOL
|
||||
return ti.decode_content
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("field_type", "number", "wire_type", "accessor"),
|
||||
[
|
||||
(UINT32_T, 2, 0, "value.as_varint()"),
|
||||
(BOOL_T, 3, 0, "value.as_varint() != 0"),
|
||||
(STRING_T, 1, 2, "value.data()"),
|
||||
(UINT32, 2, 0, "value.as_varint()"),
|
||||
(BOOL, 3, 0, "value.as_varint() != 0"),
|
||||
(STRING, 1, 2, "value.data()"),
|
||||
(FLOAT, 4, 5, "value.as_float()"),
|
||||
(FIXED32, 5, 5, "value.as_fixed32()"),
|
||||
],
|
||||
@@ -222,21 +209,18 @@ def test_decode_cases_carry_field_number_and_wire_type(
|
||||
field_type: int, number: int, wire_type: int, accessor: str
|
||||
) -> None:
|
||||
"""Each decoded field yields one case keyed on its number and declared wire type."""
|
||||
cases = _decode_cases(field_type, number)
|
||||
assert len(cases) == 1, cases
|
||||
lines = cases[0].splitlines()
|
||||
assert lines[0] == f"case PROTO_DECODE_CASE({number}, {wire_type}):", cases[0]
|
||||
assert lines[1].strip() == f"PROTO_DECODE_GUARD(tag, {number}, {wire_type});", (
|
||||
cases[0]
|
||||
)
|
||||
assert accessor in cases[0], cases[0]
|
||||
case = _decode_case(field_type, number)
|
||||
lines = case.splitlines()
|
||||
assert lines[0] == f"case PROTO_DECODE_CASE({number}, {wire_type}):", case
|
||||
assert lines[1].strip() == f"PROTO_DECODE_GUARD(tag, {number}, {wire_type});", case
|
||||
assert accessor in case, case
|
||||
|
||||
|
||||
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_T)
|
||||
desc.field.add(name="count", number=2, type=UINT32_T)
|
||||
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 = "bool decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) override;"
|
||||
@@ -256,10 +240,8 @@ def test_message_gets_a_single_decode_field_override() -> None:
|
||||
|
||||
def test_multi_statement_decode_cases_are_scoped() -> None:
|
||||
"""Bodies with several statements or locals get their own block so no jump crosses an initialization."""
|
||||
bytes_type = descriptor_pb2.FieldDescriptorProto.TYPE_BYTES
|
||||
cases = _decode_cases(bytes_type, 4)
|
||||
assert len(cases) == 1, cases
|
||||
lines = cases[0].splitlines()
|
||||
assert lines[0] == "case PROTO_DECODE_CASE(4, 2): {", cases[0]
|
||||
assert lines[-1] == "}", cases[0]
|
||||
assert "value.data();" in cases[0] and "value.size();" in cases[0]
|
||||
case = _decode_case(BYTES, 4)
|
||||
lines = case.splitlines()
|
||||
assert lines[0] == "case PROTO_DECODE_CASE(4, 2): {", case
|
||||
assert lines[-1] == "}", case
|
||||
assert "value.data();" in case and "value.size();" in case
|
||||
|
||||
Reference in New Issue
Block a user