_from_binary.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386
  1. # Copyright (c) 2025-2026 Buf Technologies, Inc.
  2. #
  3. # Licensed under the Apache License, Version 2.0 (the "License");
  4. # you may not use this file except in compliance with the License.
  5. # You may obtain a copy of the License at
  6. #
  7. # http://www.apache.org/licenses/LICENSE-2.0
  8. #
  9. # Unless required by applicable law or agreed to in writing, software
  10. # distributed under the License is distributed on an "AS IS" BASIS,
  11. # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
  12. # See the License for the specific language governing permissions and
  13. # limitations under the License.
  14. from __future__ import annotations
  15. from dataclasses import dataclass
  16. from typing import TYPE_CHECKING, Any, TypeVar
  17. from ._descriptors import (
  18. DescEnum,
  19. DescFieldValueEnum,
  20. DescFieldValueList,
  21. DescFieldValueMap,
  22. DescFieldValueMessage,
  23. DescFieldValueScalar,
  24. DescMessage,
  25. ScalarType,
  26. )
  27. from ._enum import Enum
  28. from ._field_values import scalar_zero_value
  29. from ._typing import assert_never
  30. from ._wire._binary_reader import DEPTH_LIMIT, BinaryReader
  31. from ._wire._binary_writer import BinaryWriter
  32. from ._wire._wire_type import WireType
  33. if TYPE_CHECKING:
  34. from ._message import Message
  35. T = TypeVar("T", bound=Message)
  36. @dataclass(slots=True, frozen=True)
  37. class FromBinaryOptions:
  38. """Options to control the behavior of from_binary.
  39. Args:
  40. ignore_unknown_fields: If `True`, unknown fields are ignored instead of being added to the message.
  41. """
  42. ignore_unknown_fields: bool = False
  43. # Dispatch table for reading scalar values. CPython currently does not generate
  44. # jump tables for ``match`` statements, and it is still fairly simple to use this
  45. # table instead.
  46. # https://github.com/python/cpython/issues/88449
  47. _SCALAR_READERS = (
  48. None, # 0: unused
  49. BinaryReader.double, # 1: DOUBLE
  50. BinaryReader.float_, # 2: FLOAT
  51. BinaryReader.int64, # 3: INT64
  52. BinaryReader.uint64, # 4: UINT64
  53. BinaryReader.int32, # 5: INT32
  54. BinaryReader.fixed64, # 6: FIXED64
  55. BinaryReader.fixed32, # 7: FIXED32
  56. BinaryReader.bool_, # 8: BOOL
  57. BinaryReader.string, # 9: STRING
  58. None, # 10: GROUP
  59. None, # 11: MESSAGE
  60. BinaryReader.bytes_, # 12: BYTES
  61. BinaryReader.uint32, # 13: UINT32
  62. None, # 14: ENUM
  63. BinaryReader.sfixed32, # 15: SFIXED32
  64. BinaryReader.sfixed64, # 16: SFIXED64
  65. BinaryReader.sint32, # 17: SINT32
  66. BinaryReader.sint64, # 18: SINT64
  67. )
  68. def read_scalar(scalar_type: ScalarType, reader: BinaryReader) -> Any:
  69. reader_method = _SCALAR_READERS[scalar_type.value]
  70. assert reader_method is not None # noqa: S101
  71. return reader_method(reader)
  72. # TODO delete this, and either:
  73. # - call the method from to_binary once implemented
  74. # - use a different representation for unknown fields
  75. def _encode_varint(value: int) -> bytes:
  76. """Encode an integer as a varint."""
  77. result = bytearray()
  78. while value > 0x7F:
  79. result.append((value & 0x7F) | 0x80)
  80. value >>= 7
  81. result.append(value)
  82. return bytes(result)
  83. def read_message(
  84. message: Message,
  85. reader: BinaryReader,
  86. opts: FromBinaryOptions,
  87. depth: int,
  88. *,
  89. length: int = 0,
  90. group_number: int | None = None,
  91. ) -> Message:
  92. if depth > DEPTH_LIMIT:
  93. msg = f"exceeded maximum recursion depth {DEPTH_LIMIT} while parsing message"
  94. raise RecursionError(msg)
  95. desc_message = message._desc
  96. end = reader.offset + length # Only used for length-delimited messages
  97. while group_number is not None or reader.offset < end:
  98. tag = reader.tag()
  99. if group_number is not None and tag.wire_type == WireType.EGROUP:
  100. if tag.number != group_number:
  101. msg = f"mismatched group end tag: expected {group_number}, got {tag.number}"
  102. raise ValueError(msg)
  103. break
  104. desc_field = desc_message._fields_by_tag.get(tag.raw)
  105. if desc_field is None: # Unknown field
  106. field_raw = reader.skip(tag.wire_type, depth + 1, field_number=tag.number)
  107. if not opts.ignore_unknown_fields:
  108. key_raw = _encode_varint((tag.number << 3) | tag.wire_type)
  109. message._get_or_init_unknown_fields().setdefault(tag.number, []).append(
  110. key_raw + bytes(field_raw)
  111. )
  112. continue
  113. match field_value := desc_field.value:
  114. case DescFieldValueScalar():
  115. message._set_member(desc_field, read_scalar(field_value.scalar, reader))
  116. case DescFieldValueMessage(
  117. message=desc_nested_message, delimited_encoding=delimited_encoding
  118. ):
  119. existing: Message | None = message._get_member(desc_field)
  120. if existing is None:
  121. existing = desc_nested_message.type()
  122. message._set_member(desc_field, existing)
  123. if delimited_encoding:
  124. read_message(
  125. existing, reader, opts, depth + 1, group_number=tag.number
  126. )
  127. else:
  128. read_message(
  129. existing, reader, opts, depth + 1, length=reader.varint()
  130. )
  131. case DescFieldValueEnum():
  132. value = read_enum(field_value.enum, reader)
  133. if isinstance(value, Enum):
  134. message._set_member(desc_field, value)
  135. elif not opts.ignore_unknown_fields:
  136. _write_unknown_enum_field(message, desc_field.number, value)
  137. case DescFieldValueList():
  138. read_list(
  139. message,
  140. message._get_member(desc_field),
  141. desc_field.number,
  142. field_value,
  143. tag.wire_type,
  144. reader,
  145. opts,
  146. depth,
  147. )
  148. case DescFieldValueMap():
  149. entry = read_map_entry(
  150. message, desc_field.number, field_value, reader, opts, depth
  151. )
  152. if entry:
  153. key, value = entry
  154. message._get_member(desc_field)[key] = value
  155. case _:
  156. assert_never(desc_field)
  157. return message
  158. def read_list(
  159. message: Message | None,
  160. list_: list,
  161. field_number: int,
  162. field_value: DescFieldValueList,
  163. wire_type: WireType,
  164. reader: BinaryReader,
  165. opts: FromBinaryOptions,
  166. depth: int,
  167. ) -> None:
  168. element_type = field_value.element
  169. # Packed repeated field
  170. if wire_type == WireType.LENGTH_DELIMITED and field_value._packable:
  171. assert isinstance(element_type, (ScalarType, DescEnum)) # noqa: S101
  172. _read_packed_list(message, list_, field_number, element_type, reader, opts)
  173. return
  174. if wire_type != field_value._unpacked_wire_type:
  175. # Wire type doesn't match expected unpacked type, skip the field.
  176. field_bytes = reader.skip(wire_type, depth + 1, field_number=field_number)
  177. if not opts.ignore_unknown_fields and message:
  178. key_raw = _encode_varint((field_number << 3) | wire_type)
  179. message._get_or_init_unknown_fields().setdefault(field_number, []).append(
  180. key_raw + bytes(field_bytes)
  181. )
  182. return
  183. match element_type:
  184. case ScalarType():
  185. value = read_scalar(element_type, reader)
  186. case DescMessage():
  187. if field_value.delimited_encoding:
  188. value = read_message(
  189. element_type.type(),
  190. reader,
  191. opts,
  192. depth + 1,
  193. group_number=field_number,
  194. )
  195. else:
  196. value = read_message(
  197. element_type.type(), reader, opts, depth + 1, length=reader.varint()
  198. )
  199. case DescEnum():
  200. value = read_enum(element_type, reader)
  201. if not isinstance(value, Enum):
  202. if not opts.ignore_unknown_fields:
  203. _write_unknown_enum_field(message, field_number, value)
  204. return
  205. case _:
  206. assert_never(element_type)
  207. list_.append(value)
  208. def _read_packed_list(
  209. message: Message | None,
  210. list_: list,
  211. field_number: int,
  212. element_type: ScalarType | DescEnum,
  213. reader: BinaryReader,
  214. opts: FromBinaryOptions,
  215. ) -> None:
  216. length = reader.varint()
  217. end = reader.offset + length
  218. while reader.offset < end:
  219. match element_type:
  220. case ScalarType():
  221. list_.append(read_scalar(element_type, reader))
  222. case DescEnum():
  223. value = read_enum(element_type, reader)
  224. if isinstance(value, Enum):
  225. list_.append(value)
  226. elif not opts.ignore_unknown_fields:
  227. # Even for packed fields we write unknown enum values as unpacked.
  228. _write_unknown_enum_field(message, field_number, value)
  229. case _:
  230. assert_never(element_type)
  231. def read_enum(desc_enum: DescEnum, reader: BinaryReader) -> Enum | int:
  232. value = reader.int32()
  233. if not desc_enum.open and not desc_enum._values_by_number.get(value):
  234. return value
  235. return desc_enum.type(value)
  236. def _write_unknown_enum_field(
  237. message: Message | None, field_number: int, value: int
  238. ) -> None:
  239. if message is None:
  240. return
  241. writer = BinaryWriter()
  242. writer.tag(field_number, WireType.VARINT)
  243. writer.int32(value)
  244. message._get_or_init_unknown_fields().setdefault(field_number, []).append(
  245. writer.finish()
  246. )
  247. def read_map_entry(
  248. message: Message | None,
  249. field_number: int,
  250. field_value: DescFieldValueMap,
  251. reader: BinaryReader,
  252. opts: FromBinaryOptions,
  253. depth: int,
  254. ) -> tuple[Any, Any] | None:
  255. start_offset = reader.offset
  256. length = reader.varint()
  257. end = reader.offset + length
  258. key: Any = None
  259. value: Any = None
  260. while reader.offset < end:
  261. tag = reader.tag()
  262. if tag.number == 1: # key
  263. if tag.wire_type != field_value._key_wire_type:
  264. _read_unknown_map_entry(
  265. message, field_number, reader, start_offset, opts, depth
  266. )
  267. return None
  268. key = read_scalar(field_value.key, reader)
  269. elif tag.number == 2: # value
  270. if tag.wire_type != field_value._value_wire_type:
  271. _read_unknown_map_entry(
  272. message, field_number, reader, start_offset, opts, depth
  273. )
  274. return None
  275. match field_value.value:
  276. case ScalarType() as scalar_type:
  277. value = read_scalar(scalar_type, reader)
  278. case DescEnum() as desc_enum:
  279. value = read_enum(desc_enum, reader)
  280. if not isinstance(value, Enum):
  281. _read_unknown_map_entry(
  282. message, field_number, reader, start_offset, opts, depth
  283. )
  284. return None
  285. case DescMessage():
  286. value = field_value.value.type()
  287. read_message(value, reader, opts, depth + 1, length=reader.varint())
  288. case _:
  289. assert_never(field_value.value)
  290. else:
  291. reader.skip(tag.wire_type, depth + 1, field_number=tag.number)
  292. if key is None:
  293. key = scalar_zero_value(field_value.key)
  294. if value is None:
  295. match field_value.value:
  296. case ScalarType() as scalar_type:
  297. value = scalar_zero_value(scalar_type)
  298. case DescEnum() as desc_enum:
  299. value = desc_enum.type(desc_enum.values[0].number)
  300. case DescMessage():
  301. value = field_value.value.type()
  302. case _:
  303. assert_never(field_value.value)
  304. return key, value
  305. def _read_unknown_map_entry(
  306. message: Message | None,
  307. field_number: int,
  308. reader: BinaryReader,
  309. start_offset: int,
  310. opts: FromBinaryOptions,
  311. depth: int,
  312. ) -> None:
  313. # The entire map entry must be recorded to unknown fields. Simplest way
  314. # is to reset the reader and skip.
  315. reader.seek(start_offset)
  316. entry_bytes = reader.skip(
  317. WireType.LENGTH_DELIMITED, depth, field_number=field_number
  318. )
  319. if not opts.ignore_unknown_fields and message:
  320. key_raw = _encode_varint((field_number << 3) | WireType.LENGTH_DELIMITED)
  321. message._get_or_init_unknown_fields().setdefault(field_number, []).append(
  322. key_raw + bytes(entry_bytes)
  323. )
  324. def merge_from_binary(
  325. message: Message, data: bytes, *, ignore_unknown_fields: bool = False
  326. ) -> None:
  327. """Parse serialized binary data, merging fields into an existing message.
  328. Merge rules by field kind:
  329. - Scalar and enum: the existing value is overwritten.
  330. - Message: recursively merged if already present, otherwise set.
  331. - Repeated: elements are appended.
  332. - Map: entries are added; existing keys are overwritten. Message-valued map entries are not merged.
  333. - Unknown fields: retained unless `ignore_unknown_fields` is `True`.
  334. Args:
  335. message: The message instance to merge into.
  336. data: Serialized binary protobuf data.
  337. ignore_unknown_fields: If `True`, unknown fields in the binary data are silently discarded.
  338. """
  339. message._merge_from_binary(data, ignore_unknown_fields)