mirror of
https://github.com/esphome/esphome.git
synced 2026-10-06 19:06:37 +00:00
[api] Pass the encode cursor by value through the protobuf helpers (#19014)
This commit is contained in:
@@ -59,7 +59,7 @@ static void verify_mac(uint64_t mac, size_t expected_bytes) {
|
||||
#ifdef ESPHOME_DEBUG_API
|
||||
uint8_t *proto_debug_end_ = api_buf.data() + api_buf.size();
|
||||
#endif
|
||||
ProtoEncode::encode_varint_raw_48bit(pos PROTO_ENCODE_DEBUG_ARG, mac);
|
||||
pos = ProtoEncode::encode_varint_raw_48bit(pos PROTO_ENCODE_DEBUG_ARG, mac);
|
||||
size_t new_len = pos - api_buf.data();
|
||||
|
||||
EXPECT_EQ(new_len, expected_bytes) << "mac=0x" << std::hex << mac << std::dec;
|
||||
|
||||
@@ -0,0 +1,58 @@
|
||||
esphome:
|
||||
name: api-encode-boundaries-test
|
||||
# Top-level area fills DeviceInfoResponse.suggested_area (field 16, a two-byte tag)
|
||||
area:
|
||||
id: kitchen_area
|
||||
name: Kitchen
|
||||
on_boot:
|
||||
- sensor.template.publish:
|
||||
id: zero_then_value
|
||||
state: 0.0
|
||||
|
||||
host:
|
||||
api:
|
||||
logger:
|
||||
level: DEBUG
|
||||
|
||||
sensor:
|
||||
- platform: template
|
||||
name: "Zero Then Value"
|
||||
id: zero_then_value
|
||||
# Negative int32 takes the ten byte varint path
|
||||
accuracy_decimals: -2
|
||||
update_interval: never
|
||||
|
||||
text_sensor:
|
||||
- platform: template
|
||||
name: "Long Text"
|
||||
id: long_text
|
||||
update_interval: never
|
||||
|
||||
number:
|
||||
- platform: template
|
||||
name: "Negative Number"
|
||||
optimistic: true
|
||||
min_value: -1000
|
||||
max_value: 1000
|
||||
step: 0.5
|
||||
initial_value: -123.5
|
||||
|
||||
select:
|
||||
- platform: template
|
||||
name: "Long Option Select"
|
||||
optimistic: true
|
||||
options:
|
||||
- short
|
||||
- "option-with-a-name-long-enough-that-its-length-prefix-needs-two-varint-bytes-when-the-list-entities-response-is-encoded-xxxxxxxxxx"
|
||||
initial_option: short
|
||||
|
||||
button:
|
||||
- platform: template
|
||||
name: "Publish Values"
|
||||
on_press:
|
||||
- sensor.template.publish:
|
||||
id: zero_then_value
|
||||
state: 12.5
|
||||
- text_sensor.template.publish:
|
||||
id: long_text
|
||||
state: !lambda return std::string(200, 'y');
|
||||
@@ -3,7 +3,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Callable
|
||||
from collections.abc import Awaitable, Callable
|
||||
import logging
|
||||
from typing import TypeVar
|
||||
|
||||
@@ -57,6 +57,58 @@ 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 self._waiters:
|
||||
if future.done():
|
||||
continue
|
||||
try:
|
||||
matched = predicate(state)
|
||||
except Exception as exc: # noqa: BLE001 the wait re-raises it, the callback must not die
|
||||
future.set_exception(exc)
|
||||
continue
|
||||
if matched:
|
||||
future.set_result(state)
|
||||
|
||||
def expect(
|
||||
self,
|
||||
predicate: Callable[[EntityState], bool],
|
||||
timeout: float = 5.0,
|
||||
label: str | None = None,
|
||||
) -> Awaitable[EntityState]:
|
||||
"""Arm a wait for the next state matching ``predicate`` and return the awaitable for it.
|
||||
|
||||
The wait is armed here, at call time, so it can be created before the action that produces
|
||||
the state and awaited afterwards; states seen before this call never match.
|
||||
"""
|
||||
entry = (predicate, asyncio.get_running_loop().create_future())
|
||||
self._waiters.append(entry)
|
||||
return self._wait(entry, timeout, label)
|
||||
|
||||
async def _wait(
|
||||
self,
|
||||
entry: tuple[Callable[[EntityState], bool], asyncio.Future[EntityState]],
|
||||
timeout: float,
|
||||
label: str | None,
|
||||
) -> EntityState:
|
||||
try:
|
||||
async with asyncio.timeout(timeout):
|
||||
return await entry[1]
|
||||
except TimeoutError:
|
||||
raise TimeoutError(
|
||||
f"no state matched {label or entry[0]} within {timeout}s"
|
||||
) from None
|
||||
finally:
|
||||
self._waiters.remove(entry)
|
||||
|
||||
|
||||
def find_entity[T: EntityInfo](
|
||||
entities: list[EntityInfo],
|
||||
object_id_substring: str,
|
||||
|
||||
@@ -0,0 +1,76 @@
|
||||
"""Encode paths at their branch boundaries: zero skipped float, fixed32 state, negative int32,
|
||||
length prefixes of two varint bytes and two byte field tags."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
|
||||
from aioesphomeapi import (
|
||||
NumberState,
|
||||
SelectInfo,
|
||||
SensorInfo,
|
||||
SensorState,
|
||||
TextSensorState,
|
||||
)
|
||||
import pytest
|
||||
|
||||
from .state_utils import InitialStateHelper, StateWaiter, require_entity
|
||||
from .types import APIClientConnectedFactory, RunCompiledFunction
|
||||
|
||||
LONG_OPTION = (
|
||||
"option-with-a-name-long-enough-that-its-length-prefix-needs-two-varint-bytes-"
|
||||
"when-the-list-entities-response-is-encoded-xxxxxxxxxx"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_api_encode_boundaries(
|
||||
yaml_config: str,
|
||||
run_compiled: RunCompiledFunction,
|
||||
api_client_connected: APIClientConnectedFactory,
|
||||
) -> None:
|
||||
async with run_compiled(yaml_config), api_client_connected() as client:
|
||||
device_info, (entities, _) = await asyncio.gather(
|
||||
client.device_info(), client.list_entities_services()
|
||||
)
|
||||
assert device_info.suggested_area == "Kitchen"
|
||||
|
||||
sensor = require_entity(entities, "zero_then_value", SensorInfo)
|
||||
assert sensor.accuracy_decimals == -2
|
||||
select = require_entity(entities, "long_option_select", SelectInfo)
|
||||
assert len(LONG_OPTION) >= 128
|
||||
assert select.options == ["short", LONG_OPTION]
|
||||
text = require_entity(entities, "long_text")
|
||||
number = require_entity(entities, "negative_number")
|
||||
button = require_entity(entities, "publish_values")
|
||||
|
||||
initial = InitialStateHelper(entities)
|
||||
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
|
||||
first = initial.initial_states[sensor.key]
|
||||
assert isinstance(first, SensorState)
|
||||
assert first.state == 0.0 and not first.missing_state
|
||||
first_number = initial.initial_states[number.key]
|
||||
assert isinstance(first_number, NumberState)
|
||||
assert first_number.state == -123.5
|
||||
|
||||
# Arm both waits before the press so no ordering of the replies can slip past them
|
||||
sensor_seen = waiter.expect(
|
||||
lambda s: (
|
||||
isinstance(s, SensorState) and s.key == sensor.key and s.state == 12.5
|
||||
),
|
||||
label="sensor 12.5",
|
||||
)
|
||||
text_seen = waiter.expect(
|
||||
lambda s: (
|
||||
isinstance(s, TextSensorState)
|
||||
and s.key == text.key
|
||||
and s.state == "y" * 200
|
||||
),
|
||||
label="text 200 x y",
|
||||
)
|
||||
client.button_command(button.key)
|
||||
await asyncio.gather(sensor_seen, text_seen)
|
||||
@@ -380,3 +380,13 @@ def test_api_version_minor_is_at_least_15() -> None:
|
||||
"clients to see api_version >= 1.15 in HelloResponse before they will "
|
||||
"ever request it."
|
||||
)
|
||||
|
||||
|
||||
def test_generated_encode_calls_keep_the_cursor() -> None:
|
||||
"""No generated ProtoEncode call may drop the returned cursor."""
|
||||
dropped = [
|
||||
line
|
||||
for line in CPP_TEXT.splitlines()
|
||||
if "ProtoEncode::" in line and "pos = ProtoEncode::" not in line
|
||||
]
|
||||
assert not dropped, dropped[:5]
|
||||
|
||||
@@ -15,9 +15,11 @@ import pytest
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).parents[4] / "script" / "api_protobuf"))
|
||||
|
||||
import aioesphomeapi.api_options_pb2 as pb # noqa: E402
|
||||
from api_protobuf import ( # noqa: E402
|
||||
MAX_MESSAGE_ID,
|
||||
_make_ifdef_line,
|
||||
create_field_type_info,
|
||||
get_varint64_ifdef,
|
||||
validate_message_id,
|
||||
)
|
||||
@@ -43,7 +45,14 @@ UINT64 = descriptor_pb2.FieldDescriptorProto.TYPE_UINT64
|
||||
INT64 = descriptor_pb2.FieldDescriptorProto.TYPE_INT64
|
||||
SINT64 = descriptor_pb2.FieldDescriptorProto.TYPE_SINT64
|
||||
UINT32 = descriptor_pb2.FieldDescriptorProto.TYPE_UINT32
|
||||
INT32 = descriptor_pb2.FieldDescriptorProto.TYPE_INT32
|
||||
SINT32 = descriptor_pb2.FieldDescriptorProto.TYPE_SINT32
|
||||
FIXED64 = descriptor_pb2.FieldDescriptorProto.TYPE_FIXED64
|
||||
FIXED32 = descriptor_pb2.FieldDescriptorProto.TYPE_FIXED32
|
||||
FLOAT = descriptor_pb2.FieldDescriptorProto.TYPE_FLOAT
|
||||
BOOL = descriptor_pb2.FieldDescriptorProto.TYPE_BOOL
|
||||
STRING = descriptor_pb2.FieldDescriptorProto.TYPE_STRING
|
||||
BYTES = descriptor_pb2.FieldDescriptorProto.TYPE_BYTES
|
||||
|
||||
|
||||
def test_no_varint64_fields() -> None:
|
||||
@@ -107,3 +116,69 @@ def test_message_id_at_maximum_is_accepted() -> None:
|
||||
def test_message_id_above_maximum_is_rejected() -> None:
|
||||
with pytest.raises(ValueError, match="exceeds the plaintext"):
|
||||
validate_message_id(MAX_MESSAGE_ID + 1, "TooBigMessage")
|
||||
|
||||
|
||||
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
|
||||
)
|
||||
if repeated:
|
||||
field.label = descriptor_pb2.FieldDescriptorProto.LABEL_REPEATED
|
||||
if force:
|
||||
field.options.Extensions[pb.force] = True
|
||||
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 = [
|
||||
BOOL,
|
||||
UINT32,
|
||||
INT32,
|
||||
UINT64,
|
||||
INT64,
|
||||
SINT32,
|
||||
FLOAT,
|
||||
FIXED32,
|
||||
STRING,
|
||||
BYTES,
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("field_type", SCALAR_TYPES)
|
||||
def test_forced_fields_use_the_force_overload_or_raw_writes(field_type: int) -> None:
|
||||
content = _encode_field(field_type, force=True)
|
||||
assert (
|
||||
"_force(" in content
|
||||
or "write_raw_byte(" in content
|
||||
or "write_tag_and_fixed32(" in content
|
||||
), content
|
||||
|
||||
|
||||
@pytest.mark.parametrize("field_type", [FLOAT, FIXED32])
|
||||
def test_single_byte_tag_fixed32_shares_the_outlined_writer(field_type: int) -> None:
|
||||
unconditional = _encode_field(field_type, force=True)
|
||||
assert unconditional.count("write_tag_and_fixed32(pos, 13,") == 1, unconditional
|
||||
guarded = _encode_field(field_type, force=False)
|
||||
assert guarded.startswith("if ("), guarded
|
||||
assert "[[likely]]" in guarded
|
||||
assert "write_tag_and_fixed32(pos, 13," in guarded
|
||||
|
||||
|
||||
@pytest.mark.parametrize("field_type", [FLOAT, FIXED32])
|
||||
def test_multi_byte_tag_fixed32_falls_back_to_the_generic_helper(
|
||||
field_type: int,
|
||||
) -> None:
|
||||
content = _encode_field(field_type, number=16)
|
||||
assert "write_tag_and_fixed32" not in content, content
|
||||
assert content.startswith("pos = ProtoEncode::encode_"), content
|
||||
|
||||
Reference in New Issue
Block a user