1# Protocol Buffers - Google's data interchange format
2# Copyright 2008 Google Inc. All rights reserved.
3#
4# Use of this source code is governed by a BSD-style
5# license that can be found in the LICENSE file or at
6# https://developers.google.com/open-source/licenses/bsd
7"""Provides a container for DescriptorProtos."""
8
9__author__ = 'matthewtoia@google.com (Matt Toia)'
10
11from typing import Dict, Iterator, Optional
12import warnings
13
14
15class Error(Exception):
16 pass
17
18
19class DescriptorDatabaseConflictingDefinitionError(Error):
20 """Raised when a proto is added with the same name & different descriptor."""
21
22
23class DescriptorDatabase(object):
24 """A container accepting FileDescriptorProtos and maps DescriptorProtos."""
25
26 def __init__(self) -> None:
27 self._file_desc_protos_by_file: Dict[
28 str, 'descriptor_pb2.FileDescriptorProto'
29 ] = {}
30 self._file_desc_protos_by_symbol: Dict[
31 str, 'descriptor_pb2.FileDescriptorProto'
32 ] = {}
33
34 def Add(self, file_desc_proto: 'descriptor_pb2.FileDescriptorProto') -> None:
35 """Adds the FileDescriptorProto and its types to this database.
36
37 Args:
38 file_desc_proto: The FileDescriptorProto to add.
39
40 Raises:
41 DescriptorDatabaseConflictingDefinitionError: if an attempt is made to
42 add a proto with the same name but different definition than an
43 existing proto in the database.
44 """
45 proto_name = file_desc_proto.name
46 if proto_name not in self._file_desc_protos_by_file:
47 self._file_desc_protos_by_file[proto_name] = file_desc_proto
48 elif self._file_desc_protos_by_file[proto_name] != file_desc_proto:
49 raise DescriptorDatabaseConflictingDefinitionError(
50 '%s already added, but with different descriptor.' % proto_name
51 )
52 else:
53 return
54
55 # Add all the top-level descriptors to the index.
56 package = file_desc_proto.package
57 for message in file_desc_proto.message_type:
58 for name in _ExtractSymbols(message, package):
59 self._AddSymbol(name, file_desc_proto)
60 for enum in file_desc_proto.enum_type:
61 self._AddSymbol(
62 ('.'.join((package, enum.name)) if package else enum.name),
63 file_desc_proto,
64 )
65 for enum_value in enum.value:
66 self._file_desc_protos_by_symbol[
67 '.'.join((package, enum_value.name)) if package else enum_value.name
68 ] = file_desc_proto
69 for extension in file_desc_proto.extension:
70 self._AddSymbol(
71 ('.'.join((package, extension.name)) if package else extension.name),
72 file_desc_proto,
73 )
74 for service in file_desc_proto.service:
75 self._AddSymbol(
76 ('.'.join((package, service.name)) if package else service.name),
77 file_desc_proto,
78 )
79
80 def FindFileByName(self, name: str) -> 'descriptor_pb2.FileDescriptorProto':
81 """Finds the file descriptor proto by file name.
82
83 Typically the file name is a relative path ending to a .proto file. The
84 proto with the given name will have to have been added to this database
85 using the Add method or else an error will be raised.
86
87 Args:
88 name: The file name to find.
89
90 Returns:
91 The file descriptor proto matching the name.
92
93 Raises:
94 KeyError if no file by the given name was added.
95 """
96
97 return self._file_desc_protos_by_file[name]
98
99 def FindFileContainingSymbol(
100 self, symbol: str
101 ) -> 'descriptor_pb2.FileDescriptorProto':
102 """Finds the file descriptor proto containing the specified symbol.
103
104 The symbol should be a fully qualified name including the file descriptor's
105 package and any containing messages. Some examples:
106
107 'some.package.name.Message'
108 'some.package.name.Message.NestedEnum'
109 'some.package.name.Message.some_field'
110
111 The file descriptor proto containing the specified symbol must be added to
112 this database using the Add method or else an error will be raised.
113
114 Args:
115 symbol: The fully qualified symbol name.
116
117 Returns:
118 The file descriptor proto containing the symbol.
119
120 Raises:
121 KeyError if no file contains the specified symbol.
122 """
123 symbol = symbol.lstrip('.')
124 try:
125 return self._file_desc_protos_by_symbol[symbol]
126 except KeyError:
127 # Fields, enum values, and nested extensions are not in
128 # _file_desc_protos_by_symbol. Try to find the top level
129 # descriptor. Non-existent nested symbol under a valid top level
130 # descriptor can also be found. The behavior is the same with
131 # protobuf C++.
132 top_level, _, _ = symbol.rpartition('.')
133 try:
134 return self._file_desc_protos_by_symbol[top_level]
135 except KeyError:
136 # Raise the original symbol as a KeyError for better diagnostics.
137 raise KeyError(symbol)
138
139 def FindFileContainingExtension(
140 self,
141 extendee_name: str,
142 extension_number: int, # pylint: disable=unused-argument
143 ) -> Optional['descriptor_pb2.FileDescriptorProto']:
144 # TODO: implement this API.
145 return None
146
147 def FindAllExtensionNumbers(self, extendee_name: str) -> list[int]: # pylint: disable=unused-argument
148 # TODO: implement this API.
149 return []
150
151 def _AddSymbol(
152 self, name: str, file_desc_proto: 'descriptor_pb2.FileDescriptorProto'
153 ) -> None:
154 if name in self._file_desc_protos_by_symbol:
155 warn_msg = (
156 'Conflict register for file "'
157 + file_desc_proto.name
158 + '": '
159 + name
160 + ' is already defined in file "'
161 + self._file_desc_protos_by_symbol[name].name
162 + '"'
163 )
164 warnings.warn(warn_msg, RuntimeWarning)
165 self._file_desc_protos_by_symbol[name] = file_desc_proto
166
167
168def _ExtractSymbols(
169 desc_proto: 'descriptor_pb2.DescriptorProto', package: str
170) -> Iterator[str]:
171 """Pulls out all the symbols from a descriptor proto.
172
173 Args:
174 desc_proto: The proto to extract symbols from.
175 package: The package containing the descriptor type.
176
177 Yields:
178 The fully qualified name found in the descriptor.
179 """
180 message_name = package + '.' + desc_proto.name if package else desc_proto.name
181 yield message_name
182 for nested_type in desc_proto.nested_type:
183 for symbol in _ExtractSymbols(nested_type, message_name):
184 yield symbol
185 for enum_type in desc_proto.enum_type:
186 yield '.'.join((message_name, enum_type.name))