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"""A database of Python protocol buffer generated symbols.
8
9SymbolDatabase is the MessageFactory for messages generated at compile time,
10and makes it easy to create new instances of a registered type, given only the
11type's protocol buffer symbol name.
12
13Example usage::
14
15 db = symbol_database.SymbolDatabase()
16
17 # Register symbols of interest, from one or multiple files.
18 db.RegisterFileDescriptor(my_proto_pb2.DESCRIPTOR)
19 db.RegisterMessage(my_proto_pb2.MyMessage)
20 db.RegisterEnumDescriptor(my_proto_pb2.MyEnum.DESCRIPTOR)
21
22 # The database can be used as a MessageFactory, to generate types based on
23 # their name:
24 types = db.GetMessages(['my_proto.proto'])
25 my_message_instance = types['MyMessage']()
26
27 # The database's underlying descriptor pool can be queried, so it's not
28 # necessary to know a type's filename to be able to generate it:
29 filename = db.pool.FindFileContainingSymbol('MyMessage')
30 my_message_instance = db.GetMessages([filename])['MyMessage']()
31
32 # This functionality is also provided directly via a convenience method:
33 my_message_instance = db.GetSymbol('MyMessage')()
34"""
35
36import warnings
37
38from google.protobuf import descriptor_pool
39from google.protobuf import message_factory
40from google.protobuf.internal import api_implementation
41
42
43class SymbolDatabase:
44 """A database of Python generated symbols."""
45
46 # local cache of registered classes.
47 _classes = {}
48
49 def __init__(self, pool=None):
50 """Initializes a new SymbolDatabase."""
51 self.pool = pool or descriptor_pool.DescriptorPool()
52
53 def RegisterMessage(self, message):
54 """Registers the given message type in the local database.
55
56 Calls to GetSymbol() and GetMessages() will return messages registered here.
57
58 Args:
59 message: A :class:`google.protobuf.message.Message` subclass (or
60 instance); its descriptor will be registered.
61
62 Returns:
63 The provided message.
64 """
65
66 desc = message.DESCRIPTOR
67 self._classes[desc] = message
68 self.RegisterMessageDescriptor(desc)
69 return message
70
71 def RegisterMessageDescriptor(self, message_descriptor):
72 """Registers the given message descriptor in the local database.
73
74 Args:
75 message_descriptor (Descriptor): the message descriptor to add.
76 """
77 if api_implementation.Type() == 'python':
78 # pylint: disable=protected-access
79 self.pool._AddDescriptor(message_descriptor)
80
81 def RegisterEnumDescriptor(self, enum_descriptor):
82 """Registers the given enum descriptor in the local database.
83
84 Args:
85 enum_descriptor (EnumDescriptor): The enum descriptor to register.
86
87 Returns:
88 EnumDescriptor: The provided descriptor.
89 """
90 if api_implementation.Type() == 'python':
91 # pylint: disable=protected-access
92 self.pool._AddEnumDescriptor(enum_descriptor)
93 return enum_descriptor
94
95 def RegisterServiceDescriptor(self, service_descriptor):
96 """Registers the given service descriptor in the local database.
97
98 Args:
99 service_descriptor (ServiceDescriptor): the service descriptor to
100 register.
101 """
102 if api_implementation.Type() == 'python':
103 # pylint: disable=protected-access
104 self.pool._AddServiceDescriptor(service_descriptor)
105
106 def RegisterFileDescriptor(self, file_descriptor):
107 """Registers the given file descriptor in the local database.
108
109 Args:
110 file_descriptor (FileDescriptor): The file descriptor to register.
111 """
112 if api_implementation.Type() == 'python':
113 # pylint: disable=protected-access
114 self.pool._InternalAddFileDescriptor(file_descriptor)
115
116 def GetSymbol(self, symbol):
117 """Tries to find a symbol in the local database.
118
119 Currently, this method only returns message.Message instances, however, if
120 may be extended in future to support other symbol types.
121
122 Args:
123 symbol (str): a protocol buffer symbol.
124
125 Returns:
126 A Python class corresponding to the symbol.
127
128 Raises:
129 KeyError: if the symbol could not be found.
130 """
131
132 return self._classes[self.pool.FindMessageTypeByName(symbol)]
133
134 def GetMessages(self, files):
135 # TODO: Fix the differences with MessageFactory.
136 """Gets all registered messages from a specified file.
137
138 Only messages already created and registered will be returned; (this is the
139 case for imported _pb2 modules)
140 But unlike MessageFactory, this version also returns already defined nested
141 messages, but does not register any message extensions.
142
143 Args:
144 files (list[str]): The file names to extract messages from.
145
146 Returns:
147 A dictionary mapping proto names to the message classes.
148
149 Raises:
150 KeyError: if a file could not be found.
151 """
152
153 def _GetAllMessages(desc):
154 """Walk a message Descriptor and recursively yields all message names."""
155 yield desc
156 for msg_desc in desc.nested_types:
157 for nested_desc in _GetAllMessages(msg_desc):
158 yield nested_desc
159
160 result = {}
161 for file_name in files:
162 file_desc = self.pool.FindFileByName(file_name)
163 for msg_desc in file_desc.message_types_by_name.values():
164 for desc in _GetAllMessages(msg_desc):
165 try:
166 result[desc.full_name] = self._classes[desc]
167 except KeyError:
168 # This descriptor has no registered class, skip it.
169 pass
170 return result
171
172
173_DEFAULT = SymbolDatabase(pool=descriptor_pool.Default())
174
175
176def Default():
177 """Returns the default SymbolDatabase."""
178 return _DEFAULT