diff --git a/esphome/components/api/api_pb2.cpp b/esphome/components/api/api_pb2.cpp index 3594dabc32..6c5b487212 100644 --- a/esphome/components/api/api_pb2.cpp +++ b/esphome/components/api/api_pb2.cpp @@ -1133,11 +1133,12 @@ SubscribeLogsResponse::calc_size_msg(const void *self) { bool NoiseEncryptionSetKeyRequest::decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) { const ProtoFieldValue value(data, scalar); switch (PROTO_DECODE_KEY(tag)) { - case PROTO_DECODE_CASE(1, 2): + case PROTO_DECODE_CASE(1, 2): { PROTO_DECODE_GUARD(tag, 1, 2); this->key = value.data(); this->key_len = value.size(); break; + } default: return false; } @@ -1246,11 +1247,12 @@ bool HomeassistantActionResponse::decode_field(uint32_t tag, const uint8_t *data this->error_message = StringRef(reinterpret_cast(value.data()), value.size()); break; #ifdef USE_API_HOMEASSISTANT_ACTION_RESPONSES_JSON - case PROTO_DECODE_CASE(4, 2): + case PROTO_DECODE_CASE(4, 2): { PROTO_DECODE_GUARD(tag, 4, 2); this->response_data = value.data(); this->response_data_len = value.size(); break; + } #endif default: return false; @@ -1360,11 +1362,12 @@ bool GetTimeResponse::decode_field(uint32_t tag, const uint8_t *data, proto_vari PROTO_DECODE_GUARD(tag, 1, 5); this->epoch_seconds = value.as_fixed32(); break; - case PROTO_DECODE_CASE(3, 2): + case PROTO_DECODE_CASE(3, 2): { PROTO_DECODE_GUARD(tag, 3, 2); value.decode_to_message(this->parsed_timezone); this->has_parsed_timezone = true; break; + } default: return false; } @@ -1489,11 +1492,12 @@ bool ExecuteServiceRequest::decode_field(uint32_t tag, const uint8_t *data, prot PROTO_DECODE_GUARD(tag, 1, 5); this->key = value.as_fixed32(); break; - case PROTO_DECODE_CASE(2, 2): + case PROTO_DECODE_CASE(2, 2): { PROTO_DECODE_GUARD(tag, 2, 2); this->args.emplace_back(); value.decode_to_message(this->args.back()); break; + } #ifdef USE_API_USER_DEFINED_ACTION_RESPONSES case PROTO_DECODE_CASE(3, 0): PROTO_DECODE_GUARD(tag, 3, 0); @@ -2873,11 +2877,12 @@ bool BluetoothGATTWriteRequest::decode_field(uint32_t tag, const uint8_t *data, PROTO_DECODE_GUARD(tag, 3, 0); this->response = value.as_varint() != 0; break; - case PROTO_DECODE_CASE(4, 2): + case PROTO_DECODE_CASE(4, 2): { PROTO_DECODE_GUARD(tag, 4, 2); this->data = value.data(); this->data_len = value.size(); break; + } default: return false; } @@ -2910,11 +2915,12 @@ bool BluetoothGATTWriteDescriptorRequest::decode_field(uint32_t tag, const uint8 PROTO_DECODE_GUARD(tag, 2, 0); this->handle = value.as_varint(); break; - case PROTO_DECODE_CASE(3, 2): + case PROTO_DECODE_CASE(3, 2): { PROTO_DECODE_GUARD(tag, 3, 2); this->data = value.data(); this->data_len = value.size(); break; + } default: return false; } @@ -3203,11 +3209,12 @@ bool VoiceAssistantEventResponse::decode_field(uint32_t tag, const uint8_t *data PROTO_DECODE_GUARD(tag, 1, 0); this->event_type = static_cast(value.as_varint()); break; - case PROTO_DECODE_CASE(2, 2): + case PROTO_DECODE_CASE(2, 2): { PROTO_DECODE_GUARD(tag, 2, 2); this->data.emplace_back(); value.decode_to_message(this->data.back()); break; + } default: return false; } @@ -3216,20 +3223,22 @@ bool VoiceAssistantEventResponse::decode_field(uint32_t tag, const uint8_t *data bool VoiceAssistantAudio::decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) { const ProtoFieldValue value(data, scalar); switch (PROTO_DECODE_KEY(tag)) { - case PROTO_DECODE_CASE(1, 2): + case PROTO_DECODE_CASE(1, 2): { PROTO_DECODE_GUARD(tag, 1, 2); this->data = value.data(); this->data_len = value.size(); break; + } case PROTO_DECODE_CASE(2, 0): PROTO_DECODE_GUARD(tag, 2, 0); this->end = value.as_varint() != 0; break; - case PROTO_DECODE_CASE(3, 2): + case PROTO_DECODE_CASE(3, 2): { PROTO_DECODE_GUARD(tag, 3, 2); this->data2 = value.data(); this->data2_len = value.size(); break; + } default: return false; } @@ -3381,11 +3390,12 @@ bool VoiceAssistantExternalWakeWord::decode_field(uint32_t tag, const uint8_t *d bool VoiceAssistantConfigurationRequest::decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) { const ProtoFieldValue value(data, scalar); switch (PROTO_DECODE_KEY(tag)) { - case PROTO_DECODE_CASE(1, 2): + case PROTO_DECODE_CASE(1, 2): { PROTO_DECODE_GUARD(tag, 1, 2); this->external_wake_words.emplace_back(); value.decode_to_message(this->external_wake_words.back()); break; + } default: return false; } @@ -4127,11 +4137,12 @@ bool UpdateCommandRequest::decode_field(uint32_t tag, const uint8_t *data, proto bool ZWaveProxyFrame::decode_field(uint32_t tag, const uint8_t *data, proto_varint_value_t scalar) { const ProtoFieldValue value(data, scalar); switch (PROTO_DECODE_KEY(tag)) { - case PROTO_DECODE_CASE(1, 2): + case PROTO_DECODE_CASE(1, 2): { PROTO_DECODE_GUARD(tag, 1, 2); this->data = value.data(); this->data_len = value.size(); break; + } default: return false; } @@ -4160,11 +4171,12 @@ bool ZWaveProxyRequest::decode_field(uint32_t tag, const uint8_t *data, proto_va PROTO_DECODE_GUARD(tag, 1, 0); this->type = static_cast(value.as_varint()); break; - case PROTO_DECODE_CASE(2, 2): + case PROTO_DECODE_CASE(2, 2): { PROTO_DECODE_GUARD(tag, 2, 2); this->data = value.data(); this->data_len = value.size(); break; + } default: return false; } @@ -4259,12 +4271,13 @@ bool InfraredRFTransmitRawTimingsRequest::decode_field(uint32_t tag, const uint8 PROTO_DECODE_GUARD(tag, 4, 0); this->repeat_count = value.as_varint(); break; - case PROTO_DECODE_CASE(5, 2): + case PROTO_DECODE_CASE(5, 2): { PROTO_DECODE_GUARD(tag, 5, 2); this->timings_data_ = value.data(); this->timings_length_ = value.size(); this->timings_count_ = count_packed_varints(value.data(), value.size()); break; + } case PROTO_DECODE_CASE(6, 0): PROTO_DECODE_GUARD(tag, 6, 0); this->modulation = value.as_varint(); @@ -4406,11 +4419,12 @@ bool SerialProxyWriteRequest::decode_field(uint32_t tag, const uint8_t *data, pr PROTO_DECODE_GUARD(tag, 1, 0); this->instance = value.as_varint(); break; - case PROTO_DECODE_CASE(2, 2): + case PROTO_DECODE_CASE(2, 2): { PROTO_DECODE_GUARD(tag, 2, 2); this->data = value.data(); this->data_len = value.size(); break; + } default: return false; } diff --git a/script/api_protobuf/api_protobuf.py b/script/api_protobuf/api_protobuf.py index b1e6cb8ca8..6f6615fa81 100755 --- a/script/api_protobuf/api_protobuf.py +++ b/script/api_protobuf/api_protobuf.py @@ -233,14 +233,17 @@ class TypeInfo(ABC): # 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) -> str: - """Emit one decode_field() case for a field and the wire type it expects.""" - return ( - f"case PROTO_DECODE_CASE({self.number}, {int(wire_type)}):\n" - f" PROTO_DECODE_GUARD(tag, {self.number}, {int(wire_type)});\n" - f" {body}\n" - f" break;" - ) + def decode_case(self, wire_type: WireType, body: str, scoped: bool = False) -> 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. + """ + label = f"case PROTO_DECODE_CASE({self.number}, {int(wire_type)}):" + guard = f"PROTO_DECODE_GUARD(tag, {self.number}, {int(wire_type)});" + if scoped: + 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: @@ -1025,7 +1028,8 @@ class MessageType(TypeInfo): return self.decode_case( WireType.LENGTH_DELIMITED, f"value.decode_to_message(this->{self.field_name});\n" - f" this->has_{self.name} = true;", + f"this->has_{self.name} = true;", + scoped=True, ) return self.decode_case( WireType.LENGTH_DELIMITED, @@ -1229,7 +1233,8 @@ class PointerToBytesBufferType(PointerToBufferTypeBase): return self.decode_case( WireType.LENGTH_DELIMITED, f"this->{self.field_name} = value.data();\n" - f" this->{self.field_name}_len = value.size();", + f"this->{self.field_name}_len = value.size();", + scoped=True, ) def dump(self, name: str) -> str: @@ -1368,8 +1373,9 @@ class PackedBufferTypeInfo(TypeInfo): 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());", + 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 @@ -1457,12 +1463,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) + 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) @property def encode_content(self) -> str: @@ -2159,7 +2165,8 @@ class RepeatedTypeInfo(TypeInfo): 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());", + scoped=True, ) if content is None: return None diff --git a/tests/unit_tests/components/api/test_api_protobuf_generator.py b/tests/unit_tests/components/api/test_api_protobuf_generator.py index ae629ff11a..908b791a00 100644 --- a/tests/unit_tests/components/api/test_api_protobuf_generator.py +++ b/tests/unit_tests/components/api/test_api_protobuf_generator.py @@ -252,3 +252,14 @@ def test_message_gets_a_single_decode_field_override() -> None: assert "const ProtoFieldValue value(data, scalar);" in cpp for number, wire_type in ((1, 2), (2, 0), (3, 5)): assert f"case PROTO_DECODE_CASE({number}, {wire_type}):" in cpp, cpp + + +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]