_to_binary.py 6.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212
  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
  17. from ._descriptors import (
  18. DescEnum,
  19. DescFieldValueEnum,
  20. DescFieldValueList,
  21. DescFieldValueMap,
  22. DescFieldValueMessage,
  23. DescFieldValueScalar,
  24. DescMessage,
  25. ScalarType,
  26. )
  27. from ._typing import assert_never
  28. from ._wire._binary_writer import BinaryWriter
  29. from ._wire._wire_type import WireType
  30. if TYPE_CHECKING:
  31. from ._message import Message
  32. @dataclass(slots=True, frozen=True)
  33. class ToBinaryOptions:
  34. """Options to control the behavior of to_binary.
  35. Args:
  36. write_unknown_fields: If `True`, unknown fields are written to the output.
  37. """
  38. write_unknown_fields: bool
  39. def _scalar_wire_type(scalar_type: ScalarType) -> WireType:
  40. match scalar_type:
  41. case ScalarType.FIXED64 | ScalarType.SFIXED64 | ScalarType.DOUBLE:
  42. return WireType.BIT64
  43. case ScalarType.FIXED32 | ScalarType.SFIXED32 | ScalarType.FLOAT:
  44. return WireType.BIT32
  45. case ScalarType.STRING | ScalarType.BYTES:
  46. return WireType.LENGTH_DELIMITED
  47. case _:
  48. return WireType.VARINT
  49. # Dispatch table for writing scalar values. CPython currently does not generate
  50. # jump tables for ``match`` statements, and it is still fairly simple to use this
  51. # table instead.
  52. # https://github.com/python/cpython/issues/88449
  53. _SCALAR_WRITERS = (
  54. None, # 0: unused
  55. BinaryWriter.double, # 1: DOUBLE
  56. BinaryWriter.float_, # 2: FLOAT
  57. BinaryWriter.int64, # 3: INT64
  58. BinaryWriter.uint64, # 4: UINT64
  59. BinaryWriter.int32, # 5: INT32
  60. BinaryWriter.fixed64, # 6: FIXED64
  61. BinaryWriter.fixed32, # 7: FIXED32
  62. BinaryWriter.bool_, # 8: BOOL
  63. BinaryWriter.string, # 9: STRING
  64. None, # 10: GROUP
  65. None, # 11: MESSAGE
  66. BinaryWriter.bytes_, # 12: BYTES
  67. BinaryWriter.uint32, # 13: UINT32
  68. None, # 14: ENUM
  69. BinaryWriter.sfixed32, # 15: SFIXED32
  70. BinaryWriter.sfixed64, # 16: SFIXED64
  71. BinaryWriter.sint32, # 17: SINT32
  72. BinaryWriter.sint64, # 18: SINT64
  73. )
  74. def _write_scalar(scalar_type: ScalarType, value: Any, writer: BinaryWriter) -> None:
  75. writer_method = _SCALAR_WRITERS[scalar_type.value]
  76. assert writer_method is not None # noqa: S101
  77. writer_method(writer, value)
  78. def write_scalar_field(
  79. number: int, scalar_type: ScalarType, value: Any, writer: BinaryWriter
  80. ) -> None:
  81. writer.tag(number, _scalar_wire_type(scalar_type))
  82. _write_scalar(scalar_type, value, writer)
  83. def write_message_field(
  84. number: int,
  85. value: Message,
  86. writer: BinaryWriter,
  87. opts: ToBinaryOptions,
  88. *,
  89. delimited_encoding: bool,
  90. ) -> None:
  91. if delimited_encoding:
  92. writer.tag(number, WireType.SGROUP)
  93. write_message(value, writer, opts)
  94. writer.tag(number, WireType.EGROUP)
  95. else:
  96. writer.tag(number, WireType.LENGTH_DELIMITED)
  97. writer.fork()
  98. write_message(value, writer, opts)
  99. writer.join()
  100. def write_list_field(
  101. field_number: int,
  102. field_value: DescFieldValueList,
  103. value: list,
  104. writer: BinaryWriter,
  105. opts: ToBinaryOptions,
  106. ) -> None:
  107. element_type = field_value.element
  108. if field_value.packed:
  109. element_type = (
  110. element_type if isinstance(element_type, ScalarType) else ScalarType.INT32
  111. )
  112. writer.tag(field_number, WireType.LENGTH_DELIMITED)
  113. writer.fork()
  114. for v in value:
  115. _write_scalar(element_type, v, writer)
  116. writer.join()
  117. return
  118. for v in value:
  119. match element_type:
  120. case ScalarType():
  121. write_scalar_field(field_number, element_type, v, writer)
  122. case DescEnum():
  123. write_scalar_field(field_number, ScalarType.INT32, v, writer)
  124. case DescMessage():
  125. write_message_field(
  126. field_number,
  127. v,
  128. writer,
  129. opts,
  130. delimited_encoding=field_value.delimited_encoding,
  131. )
  132. case _:
  133. assert_never(element_type)
  134. def write_map_field(
  135. number: int,
  136. field_value: DescFieldValueMap,
  137. map_: dict,
  138. writer: BinaryWriter,
  139. opts: ToBinaryOptions,
  140. ) -> None:
  141. for key, value in map_.items():
  142. writer.tag(number, WireType.LENGTH_DELIMITED)
  143. writer.fork()
  144. write_scalar_field(1, field_value.key, key, writer)
  145. match field_value.value:
  146. case ScalarType() as scalar_type:
  147. write_scalar_field(2, scalar_type, value, writer)
  148. case DescEnum():
  149. write_scalar_field(2, ScalarType.INT32, value, writer)
  150. case DescMessage():
  151. write_message_field(2, value, writer, opts, delimited_encoding=False)
  152. case _:
  153. assert_never(field_value.value)
  154. writer.join()
  155. def write_message(
  156. message: Message, writer: BinaryWriter, opts: ToBinaryOptions
  157. ) -> None:
  158. for desc_field in message:
  159. value = message._get_member(desc_field)
  160. match field_value := desc_field.value:
  161. case DescFieldValueScalar():
  162. write_scalar_field(desc_field.number, field_value.scalar, value, writer)
  163. case DescFieldValueMessage():
  164. write_message_field(
  165. desc_field.number,
  166. value,
  167. writer,
  168. opts,
  169. delimited_encoding=field_value.delimited_encoding,
  170. )
  171. case DescFieldValueEnum():
  172. write_scalar_field(desc_field.number, ScalarType.INT32, value, writer)
  173. case DescFieldValueList():
  174. write_list_field(desc_field.number, field_value, value, writer, opts)
  175. case DescFieldValueMap():
  176. write_map_field(desc_field.number, field_value, value, writer, opts)
  177. case _:
  178. assert_never(field_value)
  179. # Add unknown fields
  180. if (uf := message._unknown_fields) and opts.write_unknown_fields:
  181. for field_bytes_list in uf.values():
  182. for field_bytes in field_bytes_list:
  183. writer.raw(field_bytes)