[api] Scope generated decode cases that declare locals

A case body with a declaration or several statements now gets its own
block, as the per wire type overrides had, so no jump to a later case
label crosses an initialization.
This commit is contained in:
J. Nick Koston
2026-09-07 15:57:13 +02:00
parent bcf812d62b
commit 0641dae9d2
3 changed files with 65 additions and 33 deletions
+28 -14
View File
@@ -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<const char *>(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<enums::VoiceAssistantEvent>(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<enums::ZWaveProxyRequestType>(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;
}
+26 -19
View File
@@ -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
@@ -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]