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"""Contains _ExtensionDict class to represent extensions."""
8
9from google.protobuf.descriptor import FieldDescriptor
10from google.protobuf.internal import type_checkers
11
12
13def _VerifyExtensionHandle(message, extension_handle):
14 """Verify that the given extension handle is valid."""
15
16 if not isinstance(extension_handle, FieldDescriptor):
17 raise KeyError(
18 'HasExtension() expects an extension handle, got: %s' % extension_handle
19 )
20
21 if not extension_handle.is_extension:
22 raise KeyError('"%s" is not an extension.' % extension_handle.full_name)
23
24 if not extension_handle.containing_type:
25 raise KeyError(
26 '"%s" is missing a containing_type.' % extension_handle.full_name
27 )
28
29 if extension_handle.containing_type is not message.DESCRIPTOR:
30 raise KeyError(
31 'Extension "%s" extends message type "%s", but this '
32 'message is of type "%s".'
33 % (
34 extension_handle.full_name,
35 extension_handle.containing_type.full_name,
36 message.DESCRIPTOR.full_name,
37 )
38 )
39
40
41# TODO: Unify error handling of "unknown extension" crap.
42# TODO: Support iteritems()-style iteration over all
43# extensions with the "has" bits turned on?
44class _ExtensionDict(object):
45 """Dict-like container for Extension fields on proto instances.
46
47 Note that in all cases we expect extension handles to be
48 FieldDescriptors.
49 """
50
51 def __init__(self, extended_message):
52 """Args:
53
54 extended_message: Message instance for which we are the Extensions dict.
55 """
56 self._extended_message = extended_message
57
58 def __getitem__(self, extension_handle):
59 """Returns the current value of the given extension handle."""
60
61 _VerifyExtensionHandle(self._extended_message, extension_handle)
62
63 result = self._extended_message._fields.get(extension_handle)
64 if result is not None:
65 return result
66
67 if extension_handle.is_repeated:
68 result = extension_handle._default_constructor(self._extended_message)
69 elif extension_handle.cpp_type == FieldDescriptor.CPPTYPE_MESSAGE:
70 message_type = extension_handle.message_type
71 if not hasattr(message_type, '_concrete_class'):
72 # pylint: disable=g-import-not-at-top
73 from google.protobuf import message_factory
74
75 message_factory.GetMessageClass(message_type)
76 if not hasattr(extension_handle.message_type, '_concrete_class'):
77 from google.protobuf import message_factory
78
79 message_factory.GetMessageClass(extension_handle.message_type)
80 result = extension_handle.message_type._concrete_class()
81 try:
82 result._SetListener(self._extended_message._listener_for_children)
83 except ReferenceError:
84 pass
85 else:
86 # Singular scalar -- just return the default without inserting into the
87 # dict.
88 return extension_handle.default_value
89
90 # Atomically check if another thread has preempted us and, if not, swap
91 # in the new object we just created. If someone has preempted us, we
92 # take that object and discard ours.
93 # WARNING: We are relying on setdefault() being atomic. This is true
94 # in CPython but we haven't investigated others. This warning appears
95 # in several other locations in this file.
96 if self._extended_message._frozen:
97 result._SetFrozen()
98 result = self._extended_message._fields.setdefault(extension_handle, result)
99
100 return result
101
102 def __eq__(self, other):
103 if not isinstance(other, self.__class__):
104 return False
105
106 my_fields = self._extended_message.ListFields()
107 other_fields = other._extended_message.ListFields()
108
109 # Get rid of non-extension fields.
110 my_fields = [field for field in my_fields if field.is_extension]
111 other_fields = [field for field in other_fields if field.is_extension]
112
113 return my_fields == other_fields
114
115 def __ne__(self, other):
116 return not self == other
117
118 def __len__(self):
119 fields = self._extended_message.ListFields()
120 # Get rid of non-extension fields.
121 extension_fields = [field for field in fields if field[0].is_extension]
122 return len(extension_fields)
123
124 def __hash__(self):
125 raise TypeError('unhashable object')
126
127 # Note that this is only meaningful for non-repeated, scalar extension
128 # fields. Note also that we may have to call _Modified() when we do
129 # successfully set a field this way, to set any necessary "has" bits in the
130 # ancestors of the extended message.
131 def __setitem__(self, extension_handle, value):
132 """If extension_handle specifies a non-repeated, scalar extension
133
134 field, sets the value of that field.
135 """
136
137 _VerifyExtensionHandle(self._extended_message, extension_handle)
138
139 self._extended_message._AssureWritable()
140
141 if (
142 extension_handle.is_repeated
143 or extension_handle.cpp_type == FieldDescriptor.CPPTYPE_MESSAGE
144 ):
145 raise TypeError(
146 'Cannot assign to extension "%s" because it is a repeated or '
147 'composite type.'
148 % extension_handle.full_name
149 )
150
151 # It's slightly wasteful to lookup the type checker each time,
152 # but we expect this to be a vanishingly uncommon case anyway.
153 type_checker = type_checkers.GetTypeChecker(extension_handle)
154 # pylint: disable=protected-access
155 self._extended_message._fields[extension_handle] = type_checker.CheckValue(
156 value
157 )
158 self._extended_message._Modified()
159
160 def __delitem__(self, extension_handle):
161 self._extended_message.ClearExtension(extension_handle)
162
163 def _FindExtensionByName(self, name):
164 """Tries to find a known extension with the specified name.
165
166 Args:
167 name: Extension full name.
168
169 Returns:
170 Extension field descriptor.
171 """
172 descriptor = self._extended_message.DESCRIPTOR
173 extensions = descriptor.file.pool._extensions_by_name[descriptor]
174 return extensions.get(name, None)
175
176 def _FindExtensionByNumber(self, number):
177 """Tries to find a known extension with the field number.
178
179 Args:
180 number: Extension field number.
181
182 Returns:
183 Extension field descriptor.
184 """
185 descriptor = self._extended_message.DESCRIPTOR
186 extensions = descriptor.file.pool._extensions_by_number[descriptor]
187 return extensions.get(number, None)
188
189 def __iter__(self):
190 # Return a generator over the populated extension fields
191 return (
192 f[0] for f in self._extended_message.ListFields() if f[0].is_extension
193 )
194
195 def __contains__(self, extension_handle):
196 _VerifyExtensionHandle(self._extended_message, extension_handle)
197
198 if extension_handle not in self._extended_message._fields:
199 return False
200
201 if extension_handle.is_repeated:
202 return bool(self._extended_message._fields.get(extension_handle))
203
204 if extension_handle.cpp_type == FieldDescriptor.CPPTYPE_MESSAGE:
205 value = self._extended_message._fields.get(extension_handle)
206 # pylint: disable=protected-access
207 return value is not None and value._is_present_in_parent
208
209 return True