_file_registry.py 25 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806
  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 itertools
  16. import re
  17. from bisect import bisect_right
  18. from typing import TYPE_CHECKING, TypeVar, cast
  19. from ._bootstrap import (
  20. _EDITION_PROTO2,
  21. _EDITION_PROTO3,
  22. _EDITION_UNSTABLE,
  23. _ENUM_TYPE_OPEN,
  24. _IDEMPOTENCY_UNKNOWN,
  25. _LABEL_REPEATED,
  26. _LABEL_REQUIRED,
  27. _MAXIMUM_EDITION,
  28. _MESSAGE_ENCODING_DELIMITED,
  29. _REPEATED_FIELD_ENCODING_PACKED,
  30. _TYPE_BYTES,
  31. _TYPE_ENUM,
  32. _TYPE_GROUP,
  33. _TYPE_MESSAGE,
  34. _TYPE_STRING,
  35. _feature_defaults,
  36. _FeatureKey,
  37. )
  38. from ._descriptors import (
  39. DescEnum,
  40. DescEnumValue,
  41. DescExtension,
  42. DescField,
  43. DescFieldValue,
  44. DescFieldValueEnum,
  45. DescFieldValueList,
  46. DescFieldValueMap,
  47. DescFieldValueMessage,
  48. DescFieldValueScalar,
  49. DescFile,
  50. DescMessage,
  51. DescMethod,
  52. DescOneof,
  53. DescService,
  54. ScalarType,
  55. SupportedFieldPresence,
  56. )
  57. from ._registry import Registry
  58. from ._sanitization import escape_class_name, escape_enum_attr, escape_message_attr
  59. from ._wire._text_format import (
  60. parse_text_format_enum_value,
  61. parse_text_format_scalar_value,
  62. )
  63. _SORTED_EDITION_KEYS: tuple[int, ...] = tuple(sorted(_feature_defaults))
  64. if TYPE_CHECKING:
  65. from collections.abc import Callable, Mapping
  66. from ._enum import Enum
  67. from ._extension import Extension
  68. from ._message import Message
  69. from .wkt._gen.descriptor_pb import (
  70. DescriptorProto,
  71. EnumDescriptorProto,
  72. FieldDescriptorProto,
  73. FileDescriptorProto,
  74. MethodDescriptorProto,
  75. MethodOptions,
  76. OneofDescriptorProto,
  77. ServiceDescriptorProto,
  78. )
  79. _StubMap = Mapping[str, type[Message | Enum] | Extension] | None
  80. def create_file_registry(
  81. proto: FileDescriptorProto,
  82. resolve: Callable[[str], DescFile | FileDescriptorProto | None],
  83. stubs: _StubMap = None,
  84. ) -> Registry:
  85. """Create a registry from a single FileDescriptorProto.
  86. Recursively resolves imports using *resolve*, which must return a
  87. `DescFile` or `FileDescriptorProto` for every transitive
  88. dependency.
  89. Args:
  90. proto: The file descriptor proto to build a registry for.
  91. resolve: Callback that resolves an import path to its
  92. descriptor.
  93. stubs: Optional stub map of generated symbols.
  94. Returns:
  95. A registry containing the file and all of its transitive
  96. dependencies.
  97. """
  98. seen = set[str]()
  99. reg = Registry()
  100. def recurse_deps(file: FileDescriptorProto) -> list[FileDescriptorProto]:
  101. deps: list[FileDescriptorProto] = []
  102. for path in file.dependency:
  103. if reg.file(path) is not None:
  104. continue
  105. if path in seen:
  106. continue
  107. dep = resolve(path)
  108. match dep:
  109. case None:
  110. msg = f"unable to resolve {path}, imported by {file.name}"
  111. raise ValueError(msg)
  112. case DescFile():
  113. reg.add(dep)
  114. case _:
  115. seen.add(dep.name)
  116. deps.append(dep)
  117. return [
  118. *deps,
  119. *itertools.chain.from_iterable(
  120. recurse_deps(nested_dep) for nested_dep in deps
  121. ),
  122. ]
  123. for file in reversed([proto, *recurse_deps(proto)]):
  124. add_file(reg, file, stubs)
  125. return reg
  126. def add_file(reg: Registry, proto: FileDescriptorProto, stubs: _StubMap) -> None:
  127. enums: list[DescEnum] = []
  128. messages: list[DescMessage] = []
  129. extensions: list[DescExtension] = []
  130. services: list[DescService] = []
  131. desc_file = DescFile(
  132. edition=_get_file_edition(proto),
  133. name=proto.name,
  134. dependencies=_find_file_dependencies(reg, proto),
  135. enums=enums,
  136. messages=messages,
  137. extensions=extensions,
  138. services=services,
  139. deprecated=proto.options.deprecated if proto.options else False,
  140. proto=proto,
  141. )
  142. enums.extend(
  143. _add_enum(reg, enum, desc_file, None, stubs) for enum in proto.enum_type
  144. )
  145. map_entries = _FileMapEntries()
  146. for msg_proto in proto.message_type:
  147. msg = _add_message(reg, msg_proto, desc_file, None, map_entries, stubs)
  148. if (options := msg_proto.options) is not None and options.map_entry:
  149. map_entries.add(msg)
  150. else:
  151. messages.append(msg)
  152. services.extend(
  153. _add_service(reg, svc_proto, file=desc_file) for svc_proto in proto.service
  154. )
  155. extensions.extend(
  156. _add_extension(reg, proto, desc_file, None, stubs) for proto in proto.extension
  157. )
  158. for msg in map_entries._map.values():
  159. _build_fields(reg, msg, map_entries)
  160. for msg in desc_file.messages:
  161. _build_fields(reg, msg, map_entries)
  162. _add_message_extensions(reg, msg, stubs)
  163. reg.add(desc_file)
  164. def _add_enum(
  165. reg: Registry,
  166. proto: EnumDescriptorProto,
  167. file: DescFile,
  168. parent: DescMessage | None,
  169. stubs: _StubMap,
  170. ) -> DescEnum:
  171. type_name = _build_type_name(proto, parent, file)
  172. values: list[DescEnumValue] = []
  173. values_by_number: dict[int, DescEnumValue] = {}
  174. values_by_name: dict[str, DescEnumValue] = {}
  175. stub = cast("type[Enum]", _find_stub(type_name, file, stubs))
  176. local_name = escape_class_name(proto.name)
  177. local_qualname = f"{parent._local_qualname}.{local_name}" if parent else local_name
  178. desc_enum = DescEnum(
  179. type_name=type_name,
  180. name=proto.name,
  181. _local_name=local_name,
  182. _local_qualname=local_qualname,
  183. file=file,
  184. parent=parent,
  185. open=_is_enum_open(proto, parent if parent is not None else file),
  186. values=values,
  187. _values_by_number=values_by_number,
  188. _values_by_name=values_by_name,
  189. deprecated=proto.options.deprecated if proto.options else False,
  190. proto=proto,
  191. _type=stub,
  192. )
  193. if stub:
  194. stub._desc = desc_enum
  195. stripped_prefix = f"{_pascal_to_upper_snake_case(proto.name)}_"
  196. unstripped_names = {value_proto.name for value_proto in proto.value}
  197. for value_proto in proto.value:
  198. local_name = value_proto.name
  199. if local_name.startswith(stripped_prefix) and len(local_name) > len(
  200. stripped_prefix
  201. ):
  202. stripped_local_name = local_name[len(stripped_prefix) :]
  203. # Only strip if it does not match any unstripped name and does not start with digit.
  204. if (
  205. stripped_local_name not in unstripped_names
  206. and not stripped_local_name[0].isdigit()
  207. ):
  208. local_name = stripped_local_name
  209. desc_value = DescEnumValue(
  210. name=value_proto.name,
  211. local_name=escape_enum_attr(local_name),
  212. parent=desc_enum,
  213. number=value_proto.number,
  214. deprecated=value_proto.options.deprecated if value_proto.options else False,
  215. proto=value_proto,
  216. )
  217. values.append(desc_value)
  218. values_by_number[value_proto.number] = desc_value
  219. values_by_name[value_proto.name] = desc_value
  220. reg.add(desc_enum)
  221. return desc_enum
  222. def _add_message(
  223. reg: Registry,
  224. msg_proto: DescriptorProto,
  225. file: DescFile,
  226. parent: DescMessage | None,
  227. map_entries: _FileMapEntries,
  228. stubs: _StubMap,
  229. ) -> DescMessage:
  230. type_name = _build_type_name(msg_proto, parent, file)
  231. nested_enums: list[DescEnum] = []
  232. nested_messages: list[DescMessage] = []
  233. nested_extensions: list[DescExtension] = []
  234. fields: list[DescField] = []
  235. oneofs: list[DescOneof] = []
  236. members: list[DescField | DescOneof] = []
  237. stub = (
  238. cast("type[Message]", _find_stub(type_name, file, stubs))
  239. if msg_proto.options is None or msg_proto.options.map_entry is False
  240. else None
  241. )
  242. local_name = escape_class_name(msg_proto.name)
  243. local_qualname = f"{parent._local_qualname}.{local_name}" if parent else local_name
  244. desc_message = DescMessage(
  245. type_name=type_name,
  246. name=msg_proto.name,
  247. file=file,
  248. parent=parent,
  249. fields=fields,
  250. oneofs=oneofs,
  251. members=members,
  252. nested_enums=nested_enums,
  253. nested_messages=nested_messages,
  254. nested_extensions=nested_extensions,
  255. deprecated=msg_proto.options.deprecated if msg_proto.options else False,
  256. proto=msg_proto,
  257. _local_name=local_name,
  258. _local_qualname=local_qualname,
  259. _type=stub,
  260. )
  261. if stub:
  262. stub._desc = desc_message
  263. nested_enums.extend(
  264. _add_enum(reg, proto, file, desc_message, stubs)
  265. for proto in msg_proto.enum_type
  266. )
  267. for nested_proto in msg_proto.nested_type:
  268. nested_msg = _add_message(
  269. reg, nested_proto, file, desc_message, map_entries, stubs
  270. )
  271. if nested_proto.options is not None and nested_proto.options.map_entry:
  272. map_entries.add(nested_msg)
  273. else:
  274. nested_messages.append(nested_msg)
  275. reg.add(desc_message)
  276. return desc_message
  277. def _build_fields(
  278. reg: Registry, msg: DescMessage, map_entries: _FileMapEntries
  279. ) -> None:
  280. fields = cast("list[DescField]", msg.fields)
  281. oneofs = cast("list[DescOneof]", msg.oneofs)
  282. oneof_fields: list[list[DescField]] = [[] for _ in range(len(msg.proto.oneof_decl))]
  283. oneof_fields_by_name: list[dict[str, DescField]] = [
  284. {} for _ in range(len(msg.proto.oneof_decl))
  285. ]
  286. members = cast("list[DescField | DescOneof]", msg.members)
  287. oneofs.extend(
  288. _build_oneof(oneof_proto, msg, oneof_fields[idx], oneof_fields_by_name[idx])
  289. for idx, oneof_proto in enumerate(msg.proto.oneof_decl)
  290. )
  291. for field_proto in msg.proto.field:
  292. in_oneof = not field_proto.proto3_optional and field_proto.has_field(
  293. "oneof_index"
  294. )
  295. desc_field = _build_field(
  296. reg,
  297. field_proto,
  298. msg,
  299. oneofs[field_proto.oneof_index] if in_oneof else None,
  300. map_entries,
  301. )
  302. fields.append(desc_field)
  303. if in_oneof:
  304. assert isinstance( # noqa: S101
  305. desc_field.value,
  306. (DescFieldValueScalar, DescFieldValueMessage, DescFieldValueEnum),
  307. )
  308. oneof_fields[field_proto.oneof_index].append(desc_field)
  309. oneof_fields_by_name[field_proto.oneof_index][desc_field.name] = desc_field
  310. if len(oneof_fields[field_proto.oneof_index]) == 1: # Only add once
  311. members.append(oneofs[field_proto.oneof_index])
  312. else:
  313. members.append(desc_field)
  314. oneofs[:] = [oneof for oneof in oneofs if len(oneof.fields) > 0]
  315. for nested_msg in msg.nested_messages:
  316. _build_fields(reg, nested_msg, map_entries)
  317. msg._finish_init()
  318. def _build_field_value(
  319. reg: Registry,
  320. proto: FieldDescriptorProto,
  321. parent: DescFile | DescMessage,
  322. oneof: DescOneof | None,
  323. map_entries: _FileMapEntries,
  324. ) -> DescFieldValue:
  325. if proto.label == _LABEL_REPEATED:
  326. map_entry = (
  327. map_entries.get(proto.type_name.removeprefix("."))
  328. if proto.type == _TYPE_MESSAGE
  329. else None
  330. )
  331. if map_entry:
  332. (key, value) = _find_map_entry_fields(map_entry)
  333. return DescFieldValueMap(key=key, value=value)
  334. if proto.type in (_TYPE_MESSAGE, _TYPE_GROUP):
  335. return DescFieldValueList(
  336. element=_assert(
  337. reg.message(proto.type_name.removeprefix(".")),
  338. f"could not find message: {proto.type_name}",
  339. ),
  340. packed=_is_packed_field(proto, parent),
  341. delimited_encoding=_is_delimited_encoding(proto, parent),
  342. )
  343. if proto.type == _TYPE_ENUM:
  344. field_enum = _assert(
  345. reg.enum(proto.type_name.removeprefix(".")),
  346. f"could not find enum: {proto.type_name}",
  347. )
  348. return DescFieldValueList(
  349. element=field_enum,
  350. delimited_encoding=False,
  351. packed=_is_packed_field(proto, parent),
  352. )
  353. return DescFieldValueList(
  354. delimited_encoding=False,
  355. element=ScalarType(proto.type),
  356. packed=_is_packed_field(proto, parent),
  357. )
  358. # Singular
  359. if proto.type in (_TYPE_MESSAGE, _TYPE_GROUP):
  360. return DescFieldValueMessage(
  361. message=_assert(
  362. reg.message(proto.type_name.removeprefix(".")),
  363. f"could not find message: {proto.type_name}",
  364. ),
  365. delimited_encoding=_is_delimited_encoding(proto, parent),
  366. oneof=oneof,
  367. )
  368. if proto.type == _TYPE_ENUM:
  369. field_enum = _assert(
  370. reg.enum(proto.type_name.removeprefix(".")),
  371. f"could not find enum: {proto.type_name}",
  372. )
  373. return DescFieldValueEnum(
  374. enum=field_enum,
  375. oneof=oneof,
  376. default_value=parse_text_format_enum_value(field_enum, proto.default_value)
  377. if proto.has_field("default_value")
  378. else None,
  379. )
  380. scalar = ScalarType(proto.type)
  381. return DescFieldValueScalar(
  382. oneof=oneof,
  383. default_value=parse_text_format_scalar_value(scalar, proto.default_value)
  384. if proto.has_field("default_value")
  385. else None,
  386. scalar=scalar,
  387. )
  388. def _build_field(
  389. reg: Registry,
  390. proto: FieldDescriptorProto,
  391. msg: DescMessage,
  392. oneof: DescOneof | None,
  393. map_entries: _FileMapEntries,
  394. ) -> DescField:
  395. field_value = _build_field_value(reg, proto, msg, oneof, map_entries)
  396. in_oneof = proto.has_field("oneof_index")
  397. return DescField(
  398. name=proto.name,
  399. value=field_value,
  400. parent=msg,
  401. local_name=escape_message_attr(proto.name)
  402. if not in_oneof or proto.proto3_optional
  403. else proto.name,
  404. number=proto.number,
  405. json_name=proto.json_name,
  406. deprecated=proto.options.deprecated if proto.options else False,
  407. presence=_get_field_presence(proto, msg, in_oneof=in_oneof),
  408. proto=proto,
  409. )
  410. def _add_extension(
  411. reg: Registry,
  412. proto: FieldDescriptorProto,
  413. file: DescFile,
  414. parent: DescMessage | None,
  415. stubs: _StubMap,
  416. ) -> DescExtension:
  417. """Build a DescExtension and append to parent list."""
  418. type_name = _build_type_name(proto, parent, file)
  419. extendee = _assert(
  420. reg.message(proto.extendee.removeprefix(".")),
  421. f"could not find extendee message: {proto.extendee}",
  422. )
  423. stub = cast("Extension", _find_stub(type_name, file, stubs))
  424. field_value = _build_field_value(reg, proto, file, None, _FileMapEntries())
  425. assert not isinstance(field_value, DescFieldValueMap), ( # noqa: S101
  426. "extensions cannot be map fields"
  427. )
  428. desc_ext = DescExtension(
  429. name=proto.name,
  430. value=field_value,
  431. type_name=type_name,
  432. file=file,
  433. parent=parent,
  434. extendee=extendee,
  435. number=proto.number,
  436. json_name=f"[{type_name}]",
  437. deprecated=proto.options.deprecated if proto.options else False,
  438. presence=_get_field_presence(proto, file, is_ext=True),
  439. proto=proto,
  440. _type=stub,
  441. )
  442. if stub:
  443. stub._desc = desc_ext
  444. reg.add(desc_ext)
  445. return desc_ext
  446. def _build_oneof(
  447. proto: OneofDescriptorProto,
  448. parent: DescMessage,
  449. fields: list[DescField],
  450. fields_by_name: dict[str, DescField],
  451. ) -> DescOneof:
  452. return DescOneof(
  453. name=proto.name,
  454. local_name=escape_message_attr(proto.name),
  455. parent=parent,
  456. fields=fields,
  457. proto=proto,
  458. _fields_by_name=fields_by_name,
  459. )
  460. def _add_service(
  461. reg: Registry, svc_proto: ServiceDescriptorProto, file: DescFile
  462. ) -> DescService:
  463. type_name = _build_type_name(svc_proto, None, file)
  464. methods: list[DescMethod] = []
  465. desc_service = DescService(
  466. type_name=type_name,
  467. name=svc_proto.name,
  468. file=file,
  469. methods=methods,
  470. deprecated=svc_proto.options.deprecated if svc_proto.options else False,
  471. proto=svc_proto,
  472. )
  473. for proto in svc_proto.method:
  474. method = _build_method(reg, proto, parent=desc_service)
  475. methods.append(method)
  476. reg.add(desc_service)
  477. return desc_service
  478. def _add_message_extensions(reg: Registry, msg: DescMessage, stubs: _StubMap) -> None:
  479. nested_extensions = cast("list[DescExtension]", msg.nested_extensions)
  480. nested_extensions.extend(
  481. _add_extension(reg, proto, msg.file, msg, stubs)
  482. for proto in msg.proto.extension
  483. )
  484. for nested_msg in msg.nested_messages:
  485. _add_message_extensions(reg, nested_msg, stubs)
  486. def _build_method(
  487. reg: Registry, method_proto: MethodDescriptorProto, parent: DescService
  488. ) -> DescMethod:
  489. # Determine method kind
  490. if method_proto.client_streaming and method_proto.server_streaming:
  491. method_kind = "bidi_streaming"
  492. elif method_proto.client_streaming:
  493. method_kind = "client_streaming"
  494. elif method_proto.server_streaming:
  495. method_kind = "server_streaming"
  496. else:
  497. method_kind = "unary"
  498. input_msg = reg.message(method_proto.input_type.removeprefix("."))
  499. output_msg = reg.message(method_proto.output_type.removeprefix("."))
  500. if not input_msg or not output_msg:
  501. msg = f"could not find input/output types for method {method_proto.name}"
  502. raise ValueError(msg)
  503. idempotency = _IDEMPOTENCY_UNKNOWN
  504. if method_proto.options:
  505. idempotency = method_proto.options.idempotency_level
  506. return DescMethod(
  507. name=method_proto.name,
  508. parent=parent,
  509. method_kind=method_kind,
  510. input=input_msg,
  511. output=output_msg,
  512. idempotency=cast("MethodOptions.IdempotencyLevel", idempotency),
  513. deprecated=method_proto.options.deprecated if method_proto.options else False,
  514. proto=method_proto,
  515. )
  516. def _build_type_name(
  517. proto: EnumDescriptorProto
  518. | DescriptorProto
  519. | ServiceDescriptorProto
  520. | FieldDescriptorProto,
  521. parent: DescMessage | DescService | None,
  522. file: DescFile,
  523. ) -> str:
  524. """Create a fully qualified name for a protobuf type or extension field.
  525. The fully qualified name for messages, enumerations, and services is
  526. constructed by concatenating the package name (if present), parent
  527. message names (for nested types), and the type name. We omit the leading
  528. dot added by protobuf compilers. Examples:
  529. - mypackage.MyMessage
  530. - mypackage.MyMessage.NestedMessage
  531. The fully qualified name for extension fields is constructed by
  532. concatenating the package name (if present), parent message names (for
  533. extensions declared within a message), and the field name. Examples:
  534. - mypackage.extfield
  535. - mypackage.MyMessage.extfield
  536. """
  537. if parent is not None:
  538. return f"{parent.type_name}.{proto.name}"
  539. if len(file.proto.package) > 0:
  540. return f"{file.proto.package}.{proto.name}"
  541. return proto.name
  542. def _find_file_dependencies(
  543. reg: Registry, proto: FileDescriptorProto
  544. ) -> list[DescFile]:
  545. deps: list[DescFile] = []
  546. for path in proto.dependency:
  547. dep = reg.file(path)
  548. if dep is None:
  549. msg = f"cannot find {path}, imported by {proto.name}"
  550. raise ValueError(msg)
  551. deps.append(dep)
  552. return deps
  553. def _find_stub(
  554. type_name: str, file: DescFile, stubs: Mapping[str, type | Extension] | None
  555. ) -> type | Extension | None:
  556. if stubs is None:
  557. return None
  558. return stubs.get(
  559. type_name
  560. if file.proto.package == ""
  561. else type_name.removeprefix(f"{file.proto.package}.")
  562. )
  563. _RE_UPPER_TO_LOWER = re.compile("([^_])([A-Z][a-z]+)")
  564. _RE_LOWER_TO_UPPER = re.compile("([a-z])([A-Z])")
  565. def _pascal_to_upper_snake_case(text: str) -> str:
  566. """Convert a PascalCase enum name to UPPER_SNAKE_CASE."""
  567. s1 = _RE_UPPER_TO_LOWER.sub(r"\1_\2", text)
  568. return _RE_LOWER_TO_UPPER.sub(r"\1_\2", s1).upper()
  569. def _find_map_entry_fields(
  570. map_entry: DescMessage,
  571. ) -> tuple[ScalarType, ScalarType | DescMessage | DescEnum]:
  572. key_f = next(f for f in map_entry.fields if f.number == 1)
  573. value_f = next(f for f in map_entry.fields if f.number == 2)
  574. if not isinstance(key_f.value, DescFieldValueScalar) or key_f.value.scalar in (
  575. ScalarType.BYTES,
  576. ScalarType.FLOAT,
  577. ScalarType.DOUBLE,
  578. ):
  579. msg = "invalid map key type"
  580. raise ValueError(msg)
  581. key = key_f.value.scalar
  582. match value_f.value:
  583. case DescFieldValueScalar(scalar=scalar):
  584. return (key, scalar)
  585. case DescFieldValueMessage(message=message):
  586. return (key, message)
  587. case DescFieldValueEnum(enum=enum):
  588. return (key, enum)
  589. case _:
  590. msg = "unexpected map value type"
  591. raise TypeError(msg)
  592. def _get_file_edition(proto: FileDescriptorProto) -> int:
  593. match proto.syntax:
  594. case "" | "proto2":
  595. return _EDITION_PROTO2
  596. case "proto3":
  597. return _EDITION_PROTO3
  598. case "editions":
  599. if proto.edition == _EDITION_UNSTABLE:
  600. # EDITION_UNSTABLE is a sandbox for in-development features. Collapse it
  601. # to maximum edition so test-only editions aren't leaked to users.
  602. return _MAXIMUM_EDITION
  603. if proto.edition > _MAXIMUM_EDITION:
  604. msg = f"{proto.name}: unsupported edition: {proto.edition}"
  605. raise ValueError(msg)
  606. return proto.edition
  607. case _:
  608. msg = f"{proto.name}: unsupported syntax"
  609. raise ValueError(msg)
  610. def _is_enum_open(proto: EnumDescriptorProto, parent: DescFile | DescMessage) -> bool:
  611. return _resolve_feature("enum_type", (proto, parent)) == _ENUM_TYPE_OPEN
  612. def _get_field_presence(
  613. proto: FieldDescriptorProto,
  614. parent: DescFile | DescMessage,
  615. *,
  616. in_oneof: bool = False,
  617. is_ext: bool = False,
  618. ) -> SupportedFieldPresence:
  619. if proto.label == _LABEL_REQUIRED:
  620. return SupportedFieldPresence.LEGACY_REQUIRED
  621. if proto.label == _LABEL_REPEATED:
  622. return SupportedFieldPresence.IMPLICIT
  623. if in_oneof or proto.proto3_optional or is_ext:
  624. return SupportedFieldPresence.EXPLICIT
  625. resolved = _resolve_feature("field_presence", (proto, parent))
  626. if resolved == SupportedFieldPresence.IMPLICIT and proto.type in (
  627. _TYPE_GROUP,
  628. _TYPE_MESSAGE,
  629. ):
  630. return SupportedFieldPresence.EXPLICIT
  631. return SupportedFieldPresence(resolved)
  632. def _is_packable_field(proto: FieldDescriptorProto) -> bool:
  633. return proto.type not in (_TYPE_STRING, _TYPE_BYTES, _TYPE_GROUP, _TYPE_MESSAGE)
  634. def _is_packed_field(
  635. proto: FieldDescriptorProto, parent: DescMessage | DescFile
  636. ) -> bool:
  637. if proto.label != _LABEL_REPEATED:
  638. return False
  639. if not _is_packable_field(proto):
  640. # length-delimited types cannot be packed
  641. return False
  642. if (options := proto.options) is not None and options.has_field("packed"):
  643. # prefer the field option over edition features
  644. return options.packed
  645. return (
  646. _resolve_feature("repeated_field_encoding", (proto, parent))
  647. == _REPEATED_FIELD_ENCODING_PACKED
  648. )
  649. def _is_delimited_encoding(
  650. proto: FieldDescriptorProto, parent: DescMessage | DescFile
  651. ) -> bool:
  652. if proto.type == _TYPE_GROUP:
  653. return True
  654. return (
  655. _resolve_feature("message_encoding", (proto, parent))
  656. == _MESSAGE_ENCODING_DELIMITED
  657. )
  658. def _resolve_feature(
  659. name: _FeatureKey,
  660. ref: DescFile
  661. | DescMessage
  662. | tuple[
  663. DescriptorProto | EnumDescriptorProto | FieldDescriptorProto,
  664. DescFile | DescMessage,
  665. ],
  666. ) -> int:
  667. proto = ref.proto if isinstance(ref, (DescFile, DescMessage)) else ref[0]
  668. if (options := proto.options) and (features := options.features):
  669. val = getattr(features, name)
  670. if val != 0:
  671. return val
  672. if isinstance(ref, DescMessage):
  673. return _resolve_feature(
  674. name, ref.parent if ref.parent is not None else ref.file
  675. )
  676. if isinstance(ref, DescFile):
  677. # Use the closest entry at or before ref.edition, per the FeatureSetDefaults spec.
  678. idx = bisect_right(_SORTED_EDITION_KEYS, ref.edition) - 1
  679. if idx < 0:
  680. msg = f"no feature defaults for edition {ref.edition}"
  681. raise ValueError(msg)
  682. return _feature_defaults[_SORTED_EDITION_KEYS[idx]][name]
  683. return _resolve_feature(name, ref[1])
  684. class _FileMapEntries:
  685. def __init__(self) -> None:
  686. self._map: dict[str, DescMessage] = {}
  687. def get(self, type_name: str) -> DescMessage | None:
  688. if type_name in self._map:
  689. return self._map[type_name]
  690. return None
  691. def add(self, desc: DescMessage) -> None:
  692. options = _assert(desc.proto.options, "map entry message must have options")
  693. if options.map_entry is False:
  694. msg = "invalid map entry"
  695. raise ValueError(msg)
  696. self._map[desc.type_name] = desc
  697. _AT = TypeVar("_AT")
  698. def _assert(value: _AT | None, msg: str) -> _AT:
  699. if value is None:
  700. raise ValueError(msg)
  701. return value