[api] Emit every encode call through one generator helper

_encode_call() owns the cursor assignment and the _force suffix, so
the convention lives in one place instead of at every emission site;
the fixed32 fast path is an arm of the generic encode_content keyed by
a per type value template. write_fixed32_le uses convert_little_endian
instead of its own byte order switch. The integration test shares a
StateWaiter from state_utils and leaves the disconnect to the fixture.
This commit is contained in:
J. Nick Koston
2026-09-07 11:09:07 +02:00
parent 252bf6ea6a
commit adbbda4072
5 changed files with 174 additions and 113 deletions
+110 -70
View File
@@ -131,6 +131,12 @@ def force_str(force: bool) -> str:
return str(force).lower()
def _encode_call(func: str, *args: str, force: bool = False) -> str:
"""Emit one ProtoEncode call; every helper takes the cursor and returns it advanced."""
suffix = "_force" if force else ""
return f"pos = ProtoEncode::{func}{suffix}({', '.join(('pos', *args))});"
class TypeInfo(ABC):
"""Base class for all type information."""
@@ -264,14 +270,16 @@ class TypeInfo(ABC):
# write_raw_byte(tag) + raw encode instead of the full encode_* method,
# eliminating the zero-check branch and encode_field_raw indirection.
# {value} is replaced with the actual field expression.
RAW_ENCODE_MAP: dict[str, str] = {
"encode_uint32": "pos = ProtoEncode::encode_varint_raw(pos, {value});",
"encode_uint64": "pos = ProtoEncode::encode_varint_raw_64(pos, {value});",
"encode_sint32": "pos = ProtoEncode::encode_varint_raw_short(pos, encode_zigzag32({value}));",
"encode_sint64": "pos = ProtoEncode::encode_varint_raw_64(pos, encode_zigzag64({value}));",
"encode_int64": "pos = ProtoEncode::encode_varint_raw_64(pos, static_cast<uint64_t>({value}));",
"encode_bool": "pos = ProtoEncode::write_raw_byte(pos, {value} ? 0x01 : 0x00);",
RAW_ENCODE_MAP: dict[str, tuple[str, str]] = {
"encode_uint32": ("encode_varint_raw", "{value}"),
"encode_uint64": ("encode_varint_raw_64", "{value}"),
"encode_sint32": ("encode_varint_raw_short", "encode_zigzag32({value})"),
"encode_sint64": ("encode_varint_raw_64", "encode_zigzag64({value})"),
"encode_int64": ("encode_varint_raw_64", "static_cast<uint64_t>({value})"),
"encode_bool": ("write_raw_byte", "{value} ? 0x01 : 0x00"),
}
# Fixed32 value expression for the shared tag+fixed32 writer; None for other wire types
fixed32_value_template: str | None = None
def _encode_with_precomputed_tag(self, value_expr: str) -> str | None:
"""Try to emit a precomputed-tag encode for a field.
@@ -288,12 +296,17 @@ class TypeInfo(ABC):
return None
max_val = self.max_value
# Only use RAW_ENCODE_MAP for forced fields or fields with max_value
raw_expr = None
raw = None
if self.force or max_val is not None:
raw_expr = self.RAW_ENCODE_MAP.get(self.encode_func)
if raw_expr is None:
raw = self.RAW_ENCODE_MAP.get(self.encode_func)
if raw is None:
return None
body = f"pos = ProtoEncode::write_raw_byte(pos, {tag});\n{raw_expr.format(value=value_expr)}"
func, arg = raw
body = (
_encode_call("write_raw_byte", str(tag))
+ "\n"
+ _encode_call(func, arg.format(value=value_expr))
)
if self.force:
return body
# Non-forced with max_value: inline zero-check + raw encode
@@ -314,14 +327,16 @@ class TypeInfo(ABC):
return None
# When max_len < 128, length varint is always 1 byte
len_encode = (
f"pos = ProtoEncode::write_raw_byte(pos, static_cast<uint8_t>({len_expr}));"
_encode_call("write_raw_byte", f"static_cast<uint8_t>({len_expr})")
if max_len is not None and max_len < 128
else f"pos = ProtoEncode::encode_varint_raw(pos, {len_expr});"
else _encode_call("encode_varint_raw", len_expr)
)
return (
f"pos = ProtoEncode::write_raw_byte(pos, {tag});\n"
f"{len_encode}\n"
f"pos = ProtoEncode::encode_raw(pos, {data_expr}, {len_expr});"
return "\n".join(
(
_encode_call("write_raw_byte", str(tag)),
len_encode,
_encode_call("encode_raw", data_expr, len_expr),
)
)
def _encode_fixed32_with_precomputed_tag(self, value_expr: str) -> str | None:
@@ -330,22 +345,25 @@ class TypeInfo(ABC):
if tag >= 128:
return None
if self.force:
return (
f"pos = ProtoEncode::write_tag_and_fixed32(pos, {tag}, {value_expr});"
)
return _encode_call("write_tag_and_fixed32", str(tag), value_expr)
return (
f"if (uint32_t raw = {value_expr}; raw != 0) [[likely]] {{\n"
f" pos = ProtoEncode::write_tag_and_fixed32(pos, {tag}, raw);\n"
f" {_encode_call('write_tag_and_fixed32', str(tag), 'raw')}\n"
"}"
)
@property
def encode_content(self) -> str:
if result := self._encode_with_precomputed_tag(f"this->{self.field_name}"):
value = f"this->{self.field_name}"
if result := self._encode_with_precomputed_tag(value):
return result
if self.force:
return f"pos = ProtoEncode::{self.encode_func}_force(pos, {self.number}, this->{self.field_name});"
return f"pos = ProtoEncode::{self.encode_func}(pos, {self.number}, this->{self.field_name});"
if self.fixed32_value_template is not None and (
result := self._encode_fixed32_with_precomputed_tag(
self.fixed32_value_template.format(value=value)
)
):
return result
return _encode_call(self.encode_func, str(self.number), value, force=self.force)
encode_func = None
@@ -650,13 +668,7 @@ class FloatType(FixedSizeTypeMixin, TypeInfo):
encode_func = "encode_float"
wire_type = WireType.FIXED32 # Uses wire type 5
@property
def encode_content(self) -> str:
if result := self._encode_fixed32_with_precomputed_tag(
f"float_to_raw(this->{self.field_name})"
):
return result
return super().encode_content
fixed32_value_template = "float_to_raw({value})"
def dump(self, name: str) -> str:
o = f'snprintf(buffer, sizeof(buffer), "%g", {name});\n'
@@ -724,7 +736,7 @@ class UInt64Type(VarintTypeMixin, TypeInfo):
if self.mac_address:
return {
**TypeInfo.RAW_ENCODE_MAP,
"encode_uint64": "pos = ProtoEncode::encode_varint_raw_48bit(pos, {value});",
"encode_uint64": ("encode_varint_raw_48bit", "{value}"),
}
return TypeInfo.RAW_ENCODE_MAP
@@ -792,13 +804,7 @@ class Fixed32Type(FixedSizeTypeMixin, TypeInfo):
o += "out.append(buffer);"
return o
@property
def encode_content(self) -> str:
if result := self._encode_fixed32_with_precomputed_tag(
f"this->{self.field_name}"
):
return result
return super().encode_content
fixed32_value_template = "{value}"
def get_size_calculation(self, name: str, force: bool = False) -> str:
field_id_size = self.calculate_field_id_size()
@@ -872,9 +878,12 @@ class StringType(TypeInfo):
f"this->{self.field_name}_ref_.size()",
):
return result
if self.force:
return f"pos = ProtoEncode::encode_string_force(pos, {self.number}, this->{self.field_name}_ref_);"
return f"pos = ProtoEncode::encode_string(pos, {self.number}, this->{self.field_name}_ref_);"
return _encode_call(
"encode_string",
str(self.number),
f"this->{self.field_name}_ref_",
force=self.force,
)
def dump(self, name):
# If name is 'it', this is a repeated field element - always use string
@@ -972,7 +981,9 @@ class MessageType(TypeInfo):
@property
def encode_content(self) -> str:
# Sub-message encoding needs buffer for backpatch/sync
return f"pos = ProtoEncode::{self.encode_func}(pos, buffer, {self.number}, this->{self.field_name});"
return _encode_call(
self.encode_func, "buffer", str(self.number), f"this->{self.field_name}"
)
@property
def decode_length(self) -> str:
@@ -1079,9 +1090,13 @@ class BytesType(TypeInfo):
f"this->{self.field_name}_ptr_", f"this->{self.field_name}_len_"
):
return result
if self.force:
return f"pos = ProtoEncode::encode_bytes_force(pos, {self.number}, this->{self.field_name}_ptr_, this->{self.field_name}_len_);"
return f"pos = ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}_ptr_, this->{self.field_name}_len_);"
return _encode_call(
"encode_bytes",
str(self.number),
f"this->{self.field_name}_ptr_",
f"this->{self.field_name}_len_",
force=self.force,
)
def dump(self, name: str) -> str:
ptr_dump = f"format_hex_pretty(this->{self.field_name}_ptr_, this->{self.field_name}_len_)"
@@ -1191,9 +1206,13 @@ class PointerToBytesBufferType(PointerToBufferTypeBase):
f"this->{self.field_name}", f"this->{self.field_name}_len"
):
return result
if self.force:
return f"pos = ProtoEncode::encode_bytes_force(pos, {self.number}, this->{self.field_name}, this->{self.field_name}_len);"
return f"pos = ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}, this->{self.field_name}_len);"
return _encode_call(
"encode_bytes",
str(self.number),
f"this->{self.field_name}",
f"this->{self.field_name}_len",
force=self.force,
)
@property
def decode_length_content(self) -> str | None:
@@ -1245,15 +1264,20 @@ class PointerToStringBufferType(PointerToBufferTypeBase):
if max_len is not None and max_len < 128 and self.force:
tag = self.calculate_tag()
if tag < 128:
return f"pos = ProtoEncode::encode_short_string_force(pos, {tag}, this->{self.field_name});"
return _encode_call(
"encode_short_string_force", str(tag), f"this->{self.field_name}"
)
if result := self._encode_bytes_with_precomputed_tag(
f"this->{self.field_name}.c_str()",
f"this->{self.field_name}.size()",
):
return result
if self.force:
return f"pos = ProtoEncode::encode_string_force(pos, {self.number}, this->{self.field_name});"
return f"pos = ProtoEncode::encode_string(pos, {self.number}, this->{self.field_name});"
return _encode_call(
"encode_string",
str(self.number),
f"this->{self.field_name}",
force=self.force,
)
@property
def decode_length_content(self) -> str | None:
@@ -1440,9 +1464,13 @@ class FixedArrayBytesType(TypeInfo):
f"this->{self.field_name}", f"this->{self.field_name}_len", max_len=max_len
):
return result
if self.force:
return f"pos = ProtoEncode::encode_bytes_force(pos, {self.number}, this->{self.field_name}, this->{self.field_name}_len);"
return f"pos = ProtoEncode::encode_bytes(pos, {self.number}, this->{self.field_name}, this->{self.field_name}_len);"
return _encode_call(
"encode_bytes",
str(self.number),
f"this->{self.field_name}",
f"this->{self.field_name}_len",
force=self.force,
)
def dump(self, name: str) -> str:
return f"out.append(format_hex_pretty({name}, {name}_len));"
@@ -1539,10 +1567,8 @@ class EnumType(VarintTypeMixin, TypeInfo):
@property
def encode_content(self) -> str:
value_expr = f"static_cast<uint32_t>(this->{self.field_name})"
if self.force:
return f"pos = ProtoEncode::{self.encode_func}_force(pos, {self.number}, {value_expr});"
return (
f"pos = ProtoEncode::{self.encode_func}(pos, {self.number}, {value_expr});"
return _encode_call(
self.encode_func, str(self.number), value_expr, force=self.force
)
def dump(self, name: str) -> str:
@@ -1722,9 +1748,9 @@ def _generate_inline_encode_block(
lines = []
lines.append(f"auto &sub_msg = {element};")
lines.append(f"pos = ProtoEncode::write_raw_byte(pos, {tag});")
lines.append(_encode_call("write_raw_byte", str(tag)))
lines.append("uint8_t *len_pos = pos;")
lines.append("pos = ProtoEncode::reserve_byte(pos);")
lines.append(_encode_call("reserve_byte"))
# Generate inline field encoding for each sub-message field
for field in sub_desc.field:
@@ -1796,15 +1822,22 @@ 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 f"pos = ProtoEncode::{self._ti.encode_func}_force(pos, {self.number}, static_cast<uint32_t>({element}));"
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):
if _is_inline_encode(self._ti.cpp_type):
return _generate_inline_encode_block(
self.number, self._ti.cpp_type, element
)
return f"pos = ProtoEncode::encode_sub_message(pos, buffer, {self.number}, {element});"
return f"pos = ProtoEncode::{self._ti.encode_func}_force(pos, {self.number}, {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)
@property
def cpp_type(self) -> str:
@@ -2156,11 +2189,18 @@ class RepeatedTypeInfo(TypeInfo):
def _encode_element_call(self, element: str) -> str:
"""Helper to generate encode call for a single element."""
if isinstance(self._ti, EnumType):
return f"pos = ProtoEncode::{self._ti.encode_func}_force(pos, {self.number}, static_cast<uint32_t>({element}));"
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 f"pos = ProtoEncode::encode_sub_message(pos, buffer, {self.number}, {element});"
return f"pos = ProtoEncode::{self._ti.encode_func}_force(pos, {self.number}, {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)
@property
def encode_content(self) -> str:
@@ -2169,7 +2209,7 @@ class RepeatedTypeInfo(TypeInfo):
# Special handling for const char* elements (when container_no_template contains "const char")
if "const char" in self._container_no_template:
o = f"for (const char *it : *this->{self.field_name}) {{\n"
o += f" pos = ProtoEncode::{self._ti.encode_func}_force(pos, {self.number}, it, strlen(it));\n"
o += f" {_encode_call(self._ti.encode_func, str(self.number), 'it', 'strlen(it)', force=True)}\n"
else:
o = f"for (const auto &it : *this->{self.field_name}) {{\n"
o += f" {self._encode_element_call('it')}\n"