_binary_reader.py 7.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228
  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. import struct
  16. from dataclasses import dataclass
  17. from typing import final
  18. from ._wire_type import WireType
  19. DEPTH_LIMIT = 100
  20. @dataclass(frozen=True, slots=True, init=False)
  21. class Tag:
  22. """Tag of a protobuf field, consisting of a field number and a wire type."""
  23. number: int
  24. wire_type: WireType
  25. raw: int
  26. def __init__(self, key: int) -> None:
  27. number = key >> 3
  28. wire_type = WireType(key & 0x07)
  29. if number == 0:
  30. msg = "invalid tag with field number 0"
  31. raise ValueError(msg)
  32. object.__setattr__(self, "number", number)
  33. object.__setattr__(self, "wire_type", wire_type)
  34. object.__setattr__(self, "raw", key)
  35. @final
  36. class BinaryReader:
  37. """A reader for deserializing Protocol Buffer wire format."""
  38. def __init__(self, data: memoryview) -> None:
  39. """Initialize a BinaryReader.
  40. Args:
  41. data: The underlying memoryview to read from.
  42. """
  43. self._data = data
  44. self._offset = 0
  45. @property
  46. def offset(self) -> int:
  47. """The current read position within the buffer."""
  48. return self._offset
  49. def read(self, size: int) -> memoryview:
  50. """Read a given number of bytes from the buffer."""
  51. if self._offset + size > len(self._data):
  52. msg = "unexpected end of buffer"
  53. raise EOFError(msg)
  54. result = self._data[self._offset : self._offset + size]
  55. self._offset += size
  56. return result
  57. def seek(self, offset: int) -> None:
  58. """Move the buffer position to the given offset."""
  59. if offset < 0 or offset > len(self._data):
  60. msg = "invalid seek offset"
  61. raise ValueError(msg)
  62. self._offset = offset
  63. def varint(self, limit: int = 10, last_byte_limit: int = 1) -> int:
  64. """Read a varint from the buffer."""
  65. result = 0
  66. shift = 0
  67. data = self._data
  68. offset = self._offset
  69. end = len(data)
  70. # protobuf guarantees varints are at most 10 bytes
  71. for i in range(limit):
  72. if offset >= end:
  73. msg = "unexpected end of buffer while reading varint"
  74. raise EOFError(msg)
  75. byte = data[offset]
  76. offset += 1
  77. result |= (byte & 0x7F) << shift
  78. shift += 7
  79. if byte < 0x80:
  80. if i == limit - 1 and byte > last_byte_limit:
  81. msg = "invalid varint"
  82. raise ValueError(msg)
  83. self._offset = offset
  84. return result
  85. msg = "invalid varint"
  86. raise ValueError(msg)
  87. def tag(self) -> Tag:
  88. """Read a varint and interpret it as a protobuf field tag, consisting of a field number and a wire type."""
  89. # Protobuf tags are limited to 5-byte varints
  90. key = self.varint(5, 0x0F)
  91. return Tag(key)
  92. def bool_(self) -> bool:
  93. """Read a varint and interpret it as a boolean."""
  94. return bool(self.varint())
  95. def int32(self) -> int:
  96. """Read a varint and interpret it as a signed 32-bit integer."""
  97. value = self.varint()
  98. value %= 1 << 32
  99. return value if value < (1 << 31) else value - (1 << 32)
  100. def int64(self) -> int:
  101. """Read a varint and interpret it as a signed 64-bit integer."""
  102. value = self.varint()
  103. return value if value < (1 << 63) else value - (1 << 64)
  104. def uint32(self) -> int:
  105. """Read a varint and interpret it as an unsigned 32-bit integer."""
  106. value = self.varint()
  107. return value % (1 << 32)
  108. def uint64(self) -> int:
  109. """Read a varint and interpret it as an unsigned 64-bit integer."""
  110. return self.varint()
  111. def sint32(self) -> int:
  112. """Read a varint and interpret it as a signed 32-bit integer with zigzag encoding."""
  113. value = self.varint()
  114. value = value % (1 << 32)
  115. return ((value + 1) >> 1) * (-1 if value & 1 else 1)
  116. def sint64(self) -> int:
  117. """Read a varint and interpret it as a signed 64-bit integer with zigzag encoding."""
  118. value = self.varint()
  119. return ((value + 1) >> 1) * (-1 if value & 1 else 1)
  120. def float_(self) -> float:
  121. """Read 4 bytes and interpret them as a 32-bit floating point number."""
  122. return struct.unpack("<f", self.read(4))[0]
  123. def double(self) -> float:
  124. """Read 8 bytes and interpret them as a 64-bit floating point number."""
  125. return struct.unpack("<d", self.read(8))[0]
  126. def fixed32(self) -> int:
  127. """Read 4 bytes and interpret them as an unsigned 32-bit integer."""
  128. return struct.unpack("<I", self.read(4))[0]
  129. def sfixed32(self) -> int:
  130. """Read 4 bytes and interpret them as a signed 32-bit integer."""
  131. return struct.unpack("<i", self.read(4))[0]
  132. def fixed64(self) -> int:
  133. """Read 8 bytes and interpret them as an unsigned 64-bit integer."""
  134. return struct.unpack("<Q", self.read(8))[0]
  135. def sfixed64(self) -> int:
  136. """Read 8 bytes and interpret them as a signed 64-bit integer."""
  137. return struct.unpack("<q", self.read(8))[0]
  138. def skip(self, wire_type: WireType, depth: int, *, field_number: int) -> memoryview:
  139. """Skip a field value and return a memoryview on the skipped bytes.
  140. Args:
  141. wire_type: The wire type of the field to skip.
  142. depth: The current recursion depth of the message tree.
  143. field_number: The field number. Used for delimited encoding.
  144. Returns:
  145. A [`memoryview`][] over the skipped bytes.
  146. """
  147. if depth > DEPTH_LIMIT:
  148. msg = f"exceeded maximum recursion depth {DEPTH_LIMIT} while skipping field {field_number}"
  149. raise RecursionError(msg)
  150. start = self._offset
  151. match wire_type:
  152. case WireType.VARINT:
  153. self.varint()
  154. case WireType.BIT64:
  155. self.read(8)
  156. case WireType.LENGTH_DELIMITED:
  157. length = self.varint()
  158. self.read(length)
  159. case WireType.BIT32:
  160. self.read(4)
  161. case WireType.SGROUP:
  162. while True:
  163. inner_tag = self.tag()
  164. if inner_tag.wire_type == WireType.EGROUP:
  165. if inner_tag.number != field_number:
  166. msg = f"mismatched group tag: expected {field_number}, got {inner_tag.number}"
  167. raise ValueError(msg)
  168. break
  169. nested_depth = (
  170. depth + 1 if inner_tag.wire_type == WireType.SGROUP else depth
  171. )
  172. self.skip(
  173. inner_tag.wire_type,
  174. depth=nested_depth,
  175. field_number=inner_tag.number,
  176. )
  177. case WireType.EGROUP:
  178. msg = "unexpected end group tag outside of group"
  179. raise ValueError(msg)
  180. return self._data[start : self._offset]
  181. def bytes_(self) -> bytes:
  182. """Read a length-delimited byte sequence."""
  183. length = self.varint()
  184. return bytes(self.read(length))
  185. def string(self) -> str:
  186. """Read a length-delimited byte sequence and decode it as UTF-8."""
  187. length = self.varint()
  188. view = self.read(length)
  189. return str(view, "utf-8")