_validate.py 9.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255
  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 math
  16. from typing import TYPE_CHECKING, Any
  17. from ._descriptors import (
  18. DescEnum,
  19. DescField,
  20. DescFieldValueEnum,
  21. DescFieldValueList,
  22. DescFieldValueMap,
  23. DescFieldValueMessage,
  24. DescFieldValueScalar,
  25. DescMessage,
  26. DescOneof,
  27. ScalarType,
  28. SupportedFieldPresence,
  29. )
  30. from ._oneof import Oneof
  31. from ._typing import assert_never
  32. if TYPE_CHECKING:
  33. from ._message import Message
  34. # Integer type limits (MIN is inclusive, MAX is exclusive upper bound)
  35. INT32_MIN = -(2**31)
  36. INT32_MAX = 2**31
  37. INT64_MIN = -(2**63)
  38. INT64_MAX = 2**63
  39. UINT32_MAX = 2**32
  40. UINT64_MAX = 2**64
  41. # Range of finite values representable as float32
  42. FLOAT32_MAX = 3.4028234663852886e38
  43. FLOAT32_MIN = -3.4028234663852886e38
  44. def validate(message: Message, /) -> None:
  45. """Validate a message for correctness before serialization.
  46. Checks that all field values have the correct types and that
  47. numeric values are within their valid ranges.
  48. Args:
  49. message: The message to validate.
  50. Raises:
  51. TypeError: If a field value has the wrong type.
  52. OverflowError: If a numeric value is out of range.
  53. ValueError: If an enum or oneof field has an invalid value.
  54. """
  55. # Iterate over members instead of fields to validate oneofs instead of
  56. # silently skipping them.
  57. for member in message._desc.members:
  58. if isinstance(member, DescOneof):
  59. if (value := getattr(message, member.local_name)) is not None:
  60. _validate_oneof(member, value)
  61. continue
  62. if not message._contains_member(member):
  63. if member.presence == SupportedFieldPresence.LEGACY_REQUIRED:
  64. msg = f"cannot encode {member}: required field not set"
  65. raise ValueError(msg)
  66. continue
  67. value: object = message._get_member(member)
  68. match member:
  69. case DescField():
  70. match field_value := member.value:
  71. case DescFieldValueScalar():
  72. _validate_scalar_type(field_value.scalar, value)
  73. case DescFieldValueMessage():
  74. _validate_message_value(field_value.message, value)
  75. case DescFieldValueEnum():
  76. _validate_enum(field_value.enum, value)
  77. case DescFieldValueList():
  78. if not isinstance(value, list):
  79. msg = f"expected list, got {type(value)}"
  80. raise TypeError(msg)
  81. for element in value:
  82. _validate_container_element(field_value.element, element)
  83. case DescFieldValueMap():
  84. if not isinstance(value, dict):
  85. msg = f"expected dict, got {type(value)}"
  86. raise TypeError(msg)
  87. for k, v in value.items():
  88. _validate_scalar_type(field_value.key, k)
  89. _validate_container_element(field_value.value, v)
  90. case _:
  91. assert_never(field_value)
  92. case _:
  93. assert_never(member)
  94. def _validate_oneof(desc: DescOneof, value: object) -> None:
  95. if not isinstance(value, Oneof):
  96. msg = f"{desc.parent.type_name}.{desc.name}: expected Oneof, got {type(value)}"
  97. raise TypeError(msg)
  98. if value.field not in desc._fields_by_name:
  99. msg = (
  100. f"{desc.parent.type_name}.{desc.name}: unknown oneof field '{value.field}'"
  101. )
  102. raise ValueError(msg)
  103. field = desc._fields_by_name[value.field] # ty: ignore[invalid-argument-type] # fails to narrow this to str.
  104. value = value.value
  105. match field.value:
  106. case DescFieldValueScalar(scalar=scalar):
  107. _validate_scalar_type(scalar, value)
  108. case DescFieldValueMessage(message=message):
  109. _validate_message_value(message, value)
  110. case DescFieldValueEnum(enum=enum):
  111. _validate_enum(enum, value)
  112. case _:
  113. msg = f"invalid oneof field type: {field}"
  114. raise TypeError(msg)
  115. def _validate_container_element(
  116. element_type: ScalarType | DescEnum | DescMessage, value: object
  117. ) -> None:
  118. match element_type:
  119. case ScalarType():
  120. _validate_scalar_type(element_type, value)
  121. case DescEnum():
  122. _validate_enum(element_type, value)
  123. case DescMessage():
  124. _validate_message_value(element_type, value)
  125. case _:
  126. assert_never(element_type)
  127. def _validate_message_value(desc: DescMessage, value: object) -> None:
  128. from ._message import Message # noqa: PLC0415
  129. if not isinstance(value, Message):
  130. msg = f"expected {desc.type_name!r}, got {type(value)}"
  131. raise TypeError(msg)
  132. if value._desc.type_name != desc.type_name:
  133. msg = f"expected {desc.type_name!r}, got {value._desc.type_name!r}"
  134. raise TypeError(msg)
  135. validate(value)
  136. def _validate_enum(desc: DescEnum, value: object) -> None:
  137. if not isinstance(value, int) or isinstance(value, bool):
  138. msg = f"expected int for enum {desc.type_name}, got {type(value)}"
  139. raise TypeError(msg)
  140. if not desc.open and value not in desc._values_by_number:
  141. msg = f"invalid enum value {value} for enum {desc.type_name}"
  142. raise ValueError(msg)
  143. def _validate_scalar_type(scalar_type: ScalarType, value: Any) -> None:
  144. match scalar_type:
  145. case ScalarType.BOOL:
  146. _validate_bool(value)
  147. case ScalarType.INT32 | ScalarType.SINT32 | ScalarType.SFIXED32:
  148. _validate_int32(value)
  149. case ScalarType.UINT32 | ScalarType.FIXED32:
  150. _validate_uint32(value)
  151. case ScalarType.INT64 | ScalarType.SINT64 | ScalarType.SFIXED64:
  152. _validate_int64(value)
  153. case ScalarType.UINT64 | ScalarType.FIXED64:
  154. _validate_uint64(value)
  155. case ScalarType.FLOAT:
  156. _validate_float32(value)
  157. case ScalarType.DOUBLE:
  158. _validate_float64(value)
  159. case ScalarType.STRING:
  160. _validate_string(value)
  161. case ScalarType.BYTES:
  162. _validate_bytes(value)
  163. case _:
  164. assert_never(scalar_type)
  165. def _validate_bool(value: object) -> None:
  166. if not isinstance(value, bool):
  167. msg = f"expected bool, got {type(value)}"
  168. raise TypeError(msg)
  169. def _validate_int32(value: object) -> None:
  170. if not isinstance(value, int) or isinstance(value, bool):
  171. msg = f"expected int, got {type(value)}"
  172. raise TypeError(msg)
  173. if not (INT32_MIN <= value < INT32_MAX):
  174. msg = f"value {value} out of range for int32"
  175. raise OverflowError(msg)
  176. def _validate_uint32(value: object) -> None:
  177. if not isinstance(value, int) or isinstance(value, bool):
  178. msg = f"expected int, got {type(value)}"
  179. raise TypeError(msg)
  180. if not (0 <= value < UINT32_MAX):
  181. msg = f"value {value} out of range for uint32"
  182. raise OverflowError(msg)
  183. def _validate_int64(value: object) -> None:
  184. if not isinstance(value, int) or isinstance(value, bool):
  185. msg = f"expected int, got {type(value)}"
  186. raise TypeError(msg)
  187. if not (INT64_MIN <= value < INT64_MAX):
  188. msg = f"value {value} out of range for int64"
  189. raise OverflowError(msg)
  190. def _validate_uint64(value: object) -> None:
  191. if not isinstance(value, int) or isinstance(value, bool):
  192. msg = f"expected int, got {type(value)}"
  193. raise TypeError(msg)
  194. if not (0 <= value < UINT64_MAX):
  195. msg = f"value {value} out of range for uint64"
  196. raise OverflowError(msg)
  197. def _validate_float32(value: object) -> None:
  198. if not isinstance(value, (int, float)) or isinstance(value, bool):
  199. msg = f"expected float, got {type(value)}"
  200. raise TypeError(msg)
  201. f = float(value)
  202. if math.isfinite(f) and not (FLOAT32_MIN <= f <= FLOAT32_MAX):
  203. msg = f"value {value} out of range for float"
  204. raise OverflowError(msg)
  205. def _validate_float64(value: object) -> None:
  206. if not isinstance(value, (int, float)) or isinstance(value, bool):
  207. msg = f"expected float, got {type(value)}"
  208. raise TypeError(msg)
  209. def _validate_string(value: object) -> None:
  210. if not isinstance(value, str):
  211. msg = f"expected str, got {type(value)}"
  212. raise TypeError(msg)
  213. def _validate_bytes(value: object) -> None:
  214. if not isinstance(value, (bytes, bytearray)):
  215. msg = f"expected bytes, got {type(value)}"
  216. raise TypeError(msg)