_unknown.py 5.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169
  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 typing import TYPE_CHECKING, Any
  16. from ._descriptors import (
  17. DescEnum,
  18. DescFieldValue,
  19. DescFieldValueEnum,
  20. DescFieldValueList,
  21. DescFieldValueMap,
  22. DescFieldValueMessage,
  23. DescFieldValueScalar,
  24. ScalarType,
  25. )
  26. from ._enum import Enum
  27. from ._from_binary import (
  28. FromBinaryOptions,
  29. read_enum,
  30. read_list,
  31. read_map_entry,
  32. read_message,
  33. read_scalar,
  34. )
  35. from ._to_binary import (
  36. ToBinaryOptions,
  37. write_list_field,
  38. write_map_field,
  39. write_message_field,
  40. write_scalar_field,
  41. )
  42. from ._typing import assert_never
  43. from ._wire import BinaryReader, BinaryWriter
  44. if TYPE_CHECKING:
  45. from ._message import Message
  46. def has_unknown_field(
  47. message: Message, number: int, field_value: DescFieldValue
  48. ) -> bool:
  49. # For closed enums, we need to actually decode the value to know if it is
  50. # known or not. We special case it to keep the common case simpler.
  51. if not (uf := message._unknown_fields):
  52. return False
  53. match field_value:
  54. case (
  55. DescFieldValueEnum(enum=DescEnum(open=False))
  56. | DescFieldValueList(element=DescEnum(open=False))
  57. ):
  58. value = get_unknown_field(message, number, field_value)
  59. if value is None or value == []:
  60. return False
  61. return number in uf
  62. def get_unknown_field( # noqa: RET503
  63. message: Message, number: int, field_value: DescFieldValue
  64. ) -> Any:
  65. if not (uf := message._unknown_fields) or not (binary_fields := uf.get(number)):
  66. return None
  67. opts = FromBinaryOptions()
  68. match field_value:
  69. case DescFieldValueScalar():
  70. reader = BinaryReader(memoryview(binary_fields[-1]))
  71. reader.tag()
  72. return read_scalar(field_value.scalar, reader)
  73. case DescFieldValueEnum():
  74. reader = BinaryReader(memoryview(binary_fields[-1]))
  75. reader.tag()
  76. enum_value = read_enum(field_value.enum, reader)
  77. if isinstance(enum_value, Enum):
  78. return enum_value
  79. return None
  80. case DescFieldValueMessage():
  81. message = field_value.message.type()
  82. for binary_field in binary_fields:
  83. reader = BinaryReader(memoryview(binary_field))
  84. tag = reader.tag()
  85. if field_value.delimited_encoding:
  86. read_message(
  87. message, reader, opts, depth=0, group_number=tag.number
  88. )
  89. else:
  90. read_message(message, reader, opts, depth=0, length=reader.varint())
  91. return message
  92. case DescFieldValueList():
  93. list_value: list[Any] = []
  94. for binary_field in binary_fields:
  95. reader = BinaryReader(memoryview(binary_field))
  96. while reader.offset < len(binary_field):
  97. tag = reader.tag()
  98. read_list(
  99. None,
  100. list_value,
  101. number,
  102. field_value,
  103. tag.wire_type,
  104. reader,
  105. opts,
  106. depth=0,
  107. )
  108. return list_value
  109. case DescFieldValueMap():
  110. map_value: dict[Any, Any] = {}
  111. for binary_field in binary_fields:
  112. reader = BinaryReader(memoryview(binary_field))
  113. while reader.offset < len(binary_field):
  114. reader.tag()
  115. if entry := read_map_entry(
  116. None, number, field_value, reader, opts, depth=0
  117. ):
  118. key, value_ = entry
  119. map_value[key] = value_
  120. return map_value
  121. case _:
  122. assert_never(field_value)
  123. def set_unknown_field(
  124. message: Message, number: int, field_value: DescFieldValue, value: Any
  125. ) -> None:
  126. writer = BinaryWriter()
  127. match field_value:
  128. case DescFieldValueScalar():
  129. write_scalar_field(number, field_value.scalar, value, writer)
  130. case DescFieldValueEnum():
  131. write_scalar_field(number, ScalarType.INT32, value, writer)
  132. case DescFieldValueMessage():
  133. write_message_field(
  134. number,
  135. value,
  136. writer,
  137. ToBinaryOptions(write_unknown_fields=True),
  138. delimited_encoding=field_value.delimited_encoding,
  139. )
  140. case DescFieldValueList():
  141. write_list_field(
  142. number,
  143. field_value,
  144. value,
  145. writer,
  146. ToBinaryOptions(write_unknown_fields=True),
  147. )
  148. case DescFieldValueMap():
  149. write_map_field(
  150. number,
  151. field_value,
  152. value,
  153. writer,
  154. ToBinaryOptions(write_unknown_fields=True),
  155. )
  156. case _:
  157. assert_never(field_value)
  158. binary = writer.finish()
  159. message._get_or_init_unknown_fields()[number] = [binary]