mirror of
https://github.com/esphome/esphome.git
synced 2026-09-11 15:27:33 +00:00
[api] Derive decode cases from the type's wire type
decode_case() reads wire_type instead of taking it at every call, the expression and the store statement are two small hooks that repeated fields override, and message fields build one body. No generated case declares a local any more, so the braced form and its test go; the compiler rejects a jump over a local if one ever appears. The wire type test now proves a dropped frame with an ordering marker instead of assuming it, and shares StateWaiter.
This commit is contained in:
@@ -229,41 +229,36 @@ class TypeInfo(ABC):
|
||||
def class_member(self) -> str:
|
||||
return f"{self.cpp_type} {self.field_name}{{{self.default_value}}};"
|
||||
|
||||
def decode_case(self, wire_type: WireType, body: str) -> str:
|
||||
"""Emit one decode_field() case, keyed through the PROTO_DECODE_* macros in proto.h.
|
||||
|
||||
Multi-statement bodies get a block so a case label never jumps over a local.
|
||||
"""
|
||||
label = f"case PROTO_DECODE_CASE({self.number}, {int(wire_type)}):"
|
||||
guard = f"PROTO_DECODE_GUARD(tag, {self.number}, {int(wire_type)});"
|
||||
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;")
|
||||
def decode_case(self, body: str) -> str:
|
||||
"""Emit one decode_field() case, keyed through the PROTO_DECODE_* macros in proto.h."""
|
||||
wire_type = int(self.wire_type)
|
||||
return f"case PROTO_DECODE_CASE({self.number}, {wire_type}):\n" + indent(
|
||||
f"PROTO_DECODE_GUARD(tag, {self.number}, {wire_type});\n{body}\nbreak;"
|
||||
)
|
||||
|
||||
# Decode expression per wire type; a decodable type sets exactly one.
|
||||
decode_varint = None
|
||||
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
|
||||
def _decode_expr(self) -> str | None:
|
||||
return next(
|
||||
(
|
||||
expr
|
||||
for expr in (self.decode_varint, self.decode_length, self.decode_32bit)
|
||||
if expr is not None
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
def _decode_store(self, expr: str) -> str:
|
||||
return f"this->{self.field_name} = {expr};"
|
||||
|
||||
@property
|
||||
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
|
||||
wire_type, content = expr
|
||||
return self.decode_case(wire_type, f"this->{self.field_name} = {content};")
|
||||
expr = self._decode_expr()
|
||||
return None if expr is None else self.decode_case(self._decode_store(expr))
|
||||
|
||||
# Mapping from encode_func to raw encode expression template.
|
||||
# When a forced field has a single-byte tag, the code generator emits
|
||||
@@ -1008,19 +1003,13 @@ class MessageType(TypeInfo):
|
||||
@property
|
||||
def decode_content(self) -> str:
|
||||
# Custom decode that doesn't use templates
|
||||
body = f"value.decode_to_message(this->{self.field_name});"
|
||||
if self._track_presence:
|
||||
# decode_to_message() cannot report failure, so setting the flag
|
||||
# afterwards only documents intent; a status-returning decode could
|
||||
# gate it for real without touching callers.
|
||||
return self.decode_case(
|
||||
WireType.LENGTH_DELIMITED,
|
||||
f"value.decode_to_message(this->{self.field_name});\n"
|
||||
f"this->has_{self.name} = true;",
|
||||
)
|
||||
return self.decode_case(
|
||||
WireType.LENGTH_DELIMITED,
|
||||
f"value.decode_to_message(this->{self.field_name});",
|
||||
)
|
||||
body += f"\nthis->has_{self.name} = true;"
|
||||
return self.decode_case(body)
|
||||
|
||||
def dump(self, name: str) -> str:
|
||||
return f"{name}.dump_to(out);"
|
||||
@@ -1217,7 +1206,6 @@ class PointerToBytesBufferType(PointerToBufferTypeBase):
|
||||
@property
|
||||
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();",
|
||||
)
|
||||
@@ -1282,7 +1270,6 @@ class PointerToStringBufferType(PointerToBufferTypeBase):
|
||||
@property
|
||||
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());",
|
||||
)
|
||||
|
||||
@@ -1356,7 +1343,6 @@ class PackedBufferTypeInfo(TypeInfo):
|
||||
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());",
|
||||
@@ -1447,7 +1433,6 @@ class FixedArrayBytesType(TypeInfo):
|
||||
@property
|
||||
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);",
|
||||
)
|
||||
@@ -2124,6 +2109,12 @@ class RepeatedTypeInfo(TypeInfo):
|
||||
"""
|
||||
return self._ti.wire_type
|
||||
|
||||
def _decode_expr(self) -> str | None:
|
||||
return self._ti._decode_expr()
|
||||
|
||||
def _decode_store(self, expr: str) -> str:
|
||||
return f"this->{self.field_name}.push_back({expr});"
|
||||
|
||||
@property
|
||||
def decode_content(self) -> str | None:
|
||||
# Pointer fields don't support decoding
|
||||
@@ -2132,17 +2123,10 @@ class RepeatedTypeInfo(TypeInfo):
|
||||
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());",
|
||||
f"value.decode_to_message(this->{self.field_name}.back());"
|
||||
)
|
||||
expr = self._ti.decode_expr()
|
||||
if expr is None:
|
||||
return None
|
||||
wire_type, content = expr
|
||||
return self.decode_case(
|
||||
wire_type, f"this->{self.field_name}.push_back({content});"
|
||||
)
|
||||
return super().decode_content
|
||||
|
||||
@property
|
||||
def _ti_is_bool(self) -> bool:
|
||||
|
||||
Reference in New Issue
Block a user