_registry.py 7.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222
  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 collections import defaultdict
  16. from typing import TYPE_CHECKING, TypeVar, final
  17. from ._descriptors import DescEnum, DescExtension, DescFile, DescMessage, DescService
  18. from ._extension import Extension
  19. from ._message import Message
  20. from ._typing import assert_never
  21. if TYPE_CHECKING:
  22. from collections.abc import Iterator
  23. from ._enum import Enum
  24. _Types = DescMessage | DescEnum | DescExtension | DescService
  25. _T = TypeVar("_T", bound=_Types)
  26. @final
  27. class Registry:
  28. """A set of descriptors for files, messages, enums, extensions, and services."""
  29. def __init__(
  30. self,
  31. *args: Registry
  32. | DescFile
  33. | DescMessage
  34. | DescEnum
  35. | DescExtension
  36. | DescService
  37. | Extension
  38. | type[Message | Enum],
  39. ) -> None:
  40. """Initializes a new registry with the given types, descriptors, or registries.
  41. See [add][] for more details.
  42. """
  43. self._files = dict[str, DescFile]()
  44. self._types = dict[str, _Types]()
  45. self._extendees: dict[str, dict[int, DescExtension]] = defaultdict(dict)
  46. self.add(*args)
  47. def add(
  48. self,
  49. *args: Registry
  50. | DescFile
  51. | DescMessage
  52. | DescEnum
  53. | DescExtension
  54. | DescService
  55. | Extension
  56. | type[Message | Enum],
  57. ) -> None:
  58. """Add types, descriptors, or other registries.
  59. Args:
  60. *args:
  61. Types, descriptors, or registries to add.
  62. All entries from registries are copied.
  63. In case of duplicates, the last entry wins.
  64. For DescFile, the types and descriptors from the defined files are all
  65. collected into this registry.
  66. For DescMessage, the message itself and all nested types are collected
  67. into this registry.
  68. """
  69. for arg in args:
  70. match arg:
  71. case type() | Extension():
  72. desc_or_reg = arg.desc()
  73. case _:
  74. desc_or_reg = arg
  75. match desc_or_reg:
  76. case DescEnum() | DescService():
  77. self._types[desc_or_reg.type_name] = desc_or_reg
  78. case DescMessage():
  79. self._types[desc_or_reg.type_name] = desc_or_reg
  80. for desc in (
  81. *desc_or_reg.nested_enums,
  82. *desc_or_reg.nested_messages,
  83. *desc_or_reg.nested_extensions,
  84. ):
  85. self.add(desc)
  86. case DescExtension():
  87. self._types[desc_or_reg.type_name] = desc_or_reg
  88. self._extendees[desc_or_reg.extendee.type_name][
  89. desc_or_reg.number
  90. ] = desc_or_reg
  91. case DescFile():
  92. self._files[desc_or_reg.name] = desc_or_reg
  93. for desc in (
  94. *desc_or_reg.enums,
  95. *desc_or_reg.messages,
  96. *desc_or_reg.extensions,
  97. *desc_or_reg.services,
  98. ):
  99. self.add(desc)
  100. case Registry():
  101. self._extend(desc_or_reg)
  102. case _:
  103. assert_never(desc_or_reg)
  104. def file(self, path: str) -> DescFile | None:
  105. """Look up a file descriptor by its path.
  106. Args:
  107. path: The protobuf file path as it appears in the file
  108. descriptor (e.g. `"example/foo.proto"`).
  109. Returns:
  110. The descriptor for the file, or `None` if not found.
  111. """
  112. return self._files.get(path)
  113. def message(self, type_name: str) -> DescMessage | None:
  114. """Look up a message descriptor by its fully qualified name.
  115. Args:
  116. type_name: The fully qualified name of the message.
  117. Returns:
  118. The descriptor for the message, or `None` if not found.
  119. """
  120. return self._get_type(type_name, DescMessage)
  121. def service(self, type_name: str) -> DescService | None:
  122. """Look up a service descriptor by its fully qualified name.
  123. Args:
  124. type_name: The fully qualified name of the service.
  125. Returns:
  126. The descriptor for the service, or `None` if not found.
  127. """
  128. return self._get_type(type_name, DescService)
  129. def enum(self, type_name: str) -> DescEnum | None:
  130. """Look up an enumeration descriptor by its fully qualified name.
  131. Args:
  132. type_name: The fully qualified name of the enum.
  133. Returns:
  134. The descriptor for the enum, or `None` if not found.
  135. """
  136. return self._get_type(type_name, DescEnum)
  137. def extension(self, type_name: str) -> DescExtension | None:
  138. """Look up an extension descriptor by its fully qualified name.
  139. Args:
  140. type_name: The fully qualified name of the extension.
  141. Returns:
  142. The descriptor for the extension, or `None` if not found.
  143. """
  144. # Avoid _get_type since we cant pass a type alias to it.
  145. msg = self._types.get(type_name)
  146. return msg if isinstance(msg, DescExtension) else None
  147. def extension_for(
  148. self, type_info: DescMessage | str | Message, number: int
  149. ) -> DescExtension | None:
  150. """Look up an extension by the message it extends and field number.
  151. Args:
  152. type_info: The extended message, either as a
  153. DescMessage or its fully qualified name.
  154. number: The extension field number.
  155. Returns:
  156. The descriptor for the extension, or `None` if not found.
  157. """
  158. match type_info:
  159. case DescMessage():
  160. msg_type_name = type_info.type_name
  161. case Message():
  162. msg_type_name = type_info.desc().type_name
  163. case str():
  164. msg_type_name = type_info
  165. return self._extendees[msg_type_name].get(number)
  166. def _get_type(self, type_name: str, typ: type[_T]) -> None | _T:
  167. msg = self._types.get(type_name)
  168. return msg if isinstance(msg, typ) else None
  169. def _extend(self, other: Registry) -> None:
  170. self._files |= other._files
  171. self._types |= other._types
  172. for key in other._extendees:
  173. self._extendees[key] |= other._extendees[key]
  174. def __iter__(
  175. self,
  176. ) -> Iterator[DescFile | DescMessage | DescEnum | DescExtension | DescService]:
  177. """Iterate over all descriptors in the registry.
  178. Yields files, messages, enums, extensions, and services.
  179. """
  180. yield from self._files.values()
  181. yield from self._types.values()
  182. __slots__ = "_extendees", "_files", "_types"