[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
+2 -8
View File
@@ -413,14 +413,8 @@ class ProtoEncode {
}
/// Unaligned little-endian store; __builtin_memcpy stays inline even under -fno-builtin-memcpy.
static inline void ESPHOME_ALWAYS_INLINE write_fixed32_le(uint8_t *__restrict__ pos, uint32_t value) {
#if __BYTE_ORDER__ == __ORDER_LITTLE_ENDIAN__
__builtin_memcpy(pos, &value, 4);
#else
pos[0] = static_cast<uint8_t>(value);
pos[1] = static_cast<uint8_t>(value >> 8);
pos[2] = static_cast<uint8_t>(value >> 16);
pos[3] = static_cast<uint8_t>(value >> 24);
#endif
const uint32_t le = convert_little_endian(value);
__builtin_memcpy(pos, &le, 4);
}
/// Write a precomputed tag byte + 32-bit value. Outlined on embedded: one copy beats inline stores per field.
static PROTO_OUTLINE_FOR_SIZE uint8_t *write_tag_and_fixed32(uint8_t *__restrict__ pos PROTO_ENCODE_DEBUG_PARAM,
+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"
+25
View File
@@ -57,6 +57,31 @@ async def wait_for_state(
return await asyncio.wait_for(future, timeout=timeout)
class StateWaiter:
"""Route one state subscription to any number of predicate waits."""
def __init__(self) -> None:
self._waiters: list[
tuple[Callable[[EntityState], bool], asyncio.Future[EntityState]]
] = []
def on_state(self, state: EntityState) -> None:
for predicate, future in list(self._waiters):
if not future.done() and predicate(state):
future.set_result(state)
async def expect(
self, predicate: Callable[[EntityState], bool], timeout: float = 5.0
) -> EntityState:
"""Wait for the next state matching ``predicate``."""
entry = (predicate, asyncio.get_running_loop().create_future())
self._waiters.append(entry)
try:
return await asyncio.wait_for(entry[1], timeout)
finally:
self._waiters.remove(entry)
def find_entity[T: EntityInfo](
entities: list[EntityInfo],
object_id_substring: str,
+23 -29
View File
@@ -1,12 +1,11 @@
"""Encode paths at their branch boundaries: zero skipped float, fixed32 state, negative int32,
length prefixes of two varint bytes, two byte field tags and a clean disconnect."""
length prefixes of two varint bytes and two byte field tags."""
from __future__ import annotations
import asyncio
from aioesphomeapi import (
EntityState,
NumberState,
SelectInfo,
SensorInfo,
@@ -15,7 +14,7 @@ from aioesphomeapi import (
)
import pytest
from .state_utils import InitialStateHelper, require_entity
from .state_utils import InitialStateHelper, StateWaiter, require_entity
from .types import APIClientConnectedFactory, RunCompiledFunction
LONG_OPTION = (
@@ -30,12 +29,12 @@ async def test_api_encode_boundaries(
run_compiled: RunCompiledFunction,
api_client_connected: APIClientConnectedFactory,
) -> None:
loop = asyncio.get_running_loop()
async with run_compiled(yaml_config), api_client_connected() as client:
device_info = await client.device_info()
device_info, (entities, _) = await asyncio.gather(
client.device_info(), client.list_entities_services()
)
assert device_info.suggested_area == "Kitchen"
entities, _ = await client.list_entities_services()
sensor = require_entity(entities, "zero_then_value", SensorInfo)
assert sensor.accuracy_decimals == -2
select = require_entity(entities, "long_option_select", SelectInfo)
@@ -45,25 +44,9 @@ async def test_api_encode_boundaries(
number = require_entity(entities, "negative_number")
button = require_entity(entities, "publish_values")
sensor_value: asyncio.Future[float] = loop.create_future()
text_value: asyncio.Future[str] = loop.create_future()
initial = InitialStateHelper(entities)
def on_state(state: EntityState) -> None:
if (
isinstance(state, SensorState)
and state.key == sensor.key
and not sensor_value.done()
):
sensor_value.set_result(state.state)
elif (
isinstance(state, TextSensorState)
and state.key == text.key
and not text_value.done()
):
text_value.set_result(state.state)
client.subscribe_states(initial.on_state_wrapper(on_state))
waiter = StateWaiter()
client.subscribe_states(initial.on_state_wrapper(waiter.on_state))
await initial.wait_for_initial_states()
# A float of exactly zero is skipped on the wire and must still read as 0.0, not missing
@@ -75,8 +58,19 @@ async def test_api_encode_boundaries(
assert first_number.state == -123.5
client.button_command(button.key)
assert await asyncio.wait_for(sensor_value, 5.0) == 12.5
assert await asyncio.wait_for(text_value, 5.0) == "y" * 200
# DisconnectRequest and DisconnectResponse carry no fields
await client.disconnect()
await asyncio.gather(
waiter.expect(
lambda s: (
isinstance(s, SensorState)
and s.key == sensor.key
and s.state == 12.5
)
),
waiter.expect(
lambda s: (
isinstance(s, TextSensorState)
and s.key == text.key
and s.state == "y" * 200
)
),
)
@@ -118,10 +118,9 @@ def test_message_id_above_maximum_is_rejected() -> None:
validate_message_id(MAX_MESSAGE_ID + 1, "TooBigMessage")
def _encode_field(
field_type: int, number: int = 1, force: bool = False, repeated: bool = False
) -> str:
"""Return the encode statement the generator emits for one encode-only field."""
def _field(
field_type: int, number: int = 1, *, force: bool = False, repeated: bool = False
) -> descriptor_pb2.FieldDescriptorProto:
field = descriptor_pb2.FieldDescriptorProto(
name="value", number=number, type=field_type
)
@@ -129,8 +128,17 @@ def _encode_field(
field.label = descriptor_pb2.FieldDescriptorProto.LABEL_REPEATED
if force:
field.options.Extensions[pb.force] = True
ti = create_field_type_info(field, needs_decode=False, needs_encode=True)
return ti.encode_content
return field
def _encode_field(
field_type: int, number: int = 1, force: bool = False, repeated: bool = False
) -> str:
"""Return the encode statement the generator emits for one encode-only field."""
field = _field(field_type, number, force=force, repeated=repeated)
return create_field_type_info(
field, needs_decode=False, needs_encode=True
).encode_content
SCALAR_TYPES = [