[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:
J. Nick Koston
2026-09-07 15:24:33 +02:00
parent 9de8c688e4
commit 3430194641
3 changed files with 65 additions and 128 deletions
+2 -4
View File
@@ -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; }
+45 -88
View File
@@ -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