mirror of
https://github.com/protocolbuffers/protobuf
synced 2026-08-26 02:23:14 -04:00
[py/pure python] Implement Make GetOptions() return immutable options. Raise a TypeError when options returned GetOptions() by is mutated.
PiperOrigin-RevId: 926114776
This commit is contained in:
parent
791fbe2249
commit
f5738bb5f0
4 changed files with 112 additions and 4 deletions
|
|
@ -45,7 +45,7 @@ class BaseContainer(Sequence[_T]):
|
|||
"""Base container class."""
|
||||
|
||||
# Minimizes memory usage and disallows assignment to other attributes.
|
||||
__slots__ = ['_message_listener', '_values']
|
||||
__slots__ = ['_message_listener', '_values', '_frozen']
|
||||
|
||||
def __init__(self, message_listener: Any) -> None:
|
||||
"""Args:
|
||||
|
|
@ -56,6 +56,7 @@ class BaseContainer(Sequence[_T]):
|
|||
"""
|
||||
self._message_listener = message_listener
|
||||
self._values = []
|
||||
self._frozen = False
|
||||
|
||||
@overload
|
||||
def __getitem__(self, key: int) -> _T:
|
||||
|
|
@ -83,7 +84,16 @@ class BaseContainer(Sequence[_T]):
|
|||
def __repr__(self) -> str:
|
||||
return repr(self._values)
|
||||
|
||||
def _SetFrozen(self) -> None:
|
||||
self._frozen = True
|
||||
|
||||
def _AssureWritable(self) -> 'BaseContainer[_T]':
|
||||
if self._frozen:
|
||||
raise TypeError('Container is immutable')
|
||||
return self
|
||||
|
||||
def sort(self, *args, **kwargs) -> None:
|
||||
self._AssureWritable()
|
||||
# Continue to support the old sort_function keyword argument.
|
||||
# This is expected to be a rare occurrence, so use LBYL to avoid
|
||||
# the overhead of actually catching KeyError.
|
||||
|
|
@ -92,6 +102,7 @@ class BaseContainer(Sequence[_T]):
|
|||
self._values.sort(*args, **kwargs)
|
||||
|
||||
def reverse(self) -> None:
|
||||
self._AssureWritable()
|
||||
self._values.reverse()
|
||||
|
||||
|
||||
|
|
@ -126,18 +137,21 @@ class RepeatedScalarFieldContainer(BaseContainer[_T], MutableSequence[_T]):
|
|||
|
||||
def append(self, value: _T) -> None:
|
||||
"""Appends an item to the list. Similar to list.append()."""
|
||||
self._AssureWritable()
|
||||
self._values.append(self._type_checker.CheckValue(value))
|
||||
if not self._message_listener.dirty:
|
||||
self._message_listener.Modified()
|
||||
|
||||
def insert(self, key: int, value: _T) -> None:
|
||||
"""Inserts the item at the specified position. Similar to list.insert()."""
|
||||
self._AssureWritable()
|
||||
self._values.insert(key, self._type_checker.CheckValue(value))
|
||||
if not self._message_listener.dirty:
|
||||
self._message_listener.Modified()
|
||||
|
||||
def extend(self, elem_seq: Iterable[_T]) -> None:
|
||||
"""Extends by appending the given iterable. Similar to list.extend()."""
|
||||
self._AssureWritable()
|
||||
elem_seq_iter = iter(elem_seq)
|
||||
new_values = [self._type_checker.CheckValue(elem) for elem in elem_seq_iter]
|
||||
if new_values:
|
||||
|
|
@ -152,16 +166,19 @@ class RepeatedScalarFieldContainer(BaseContainer[_T], MutableSequence[_T]):
|
|||
|
||||
one. We do not check the types of the individual fields.
|
||||
"""
|
||||
self._AssureWritable()
|
||||
self._values.extend(other)
|
||||
self._message_listener.Modified()
|
||||
|
||||
def remove(self, elem: _T):
|
||||
"""Removes an item from the list. Similar to list.remove()."""
|
||||
self._AssureWritable()
|
||||
self._values.remove(elem)
|
||||
self._message_listener.Modified()
|
||||
|
||||
def pop(self, key: Optional[int] = -1) -> _T:
|
||||
"""Removes and returns an item at a given index. Similar to list.pop()."""
|
||||
self._AssureWritable()
|
||||
value = self._values[key]
|
||||
self.__delitem__(key)
|
||||
return value
|
||||
|
|
@ -176,6 +193,7 @@ class RepeatedScalarFieldContainer(BaseContainer[_T], MutableSequence[_T]):
|
|||
|
||||
def __setitem__(self, key, value) -> None:
|
||||
"""Sets the item on the specified position."""
|
||||
self._AssureWritable()
|
||||
if isinstance(key, slice):
|
||||
if key.step is not None:
|
||||
raise ValueError('Extended slices not supported')
|
||||
|
|
@ -187,6 +205,7 @@ class RepeatedScalarFieldContainer(BaseContainer[_T], MutableSequence[_T]):
|
|||
|
||||
def __delitem__(self, key: Union[int, slice]) -> None:
|
||||
"""Deletes the item at the specified position."""
|
||||
self._AssureWritable()
|
||||
del self._values[key]
|
||||
self._message_listener.Modified()
|
||||
|
||||
|
|
@ -273,11 +292,17 @@ class RepeatedCompositeFieldContainer(BaseContainer[_T], MutableSequence[_T]):
|
|||
super().__init__(message_listener)
|
||||
self._message_descriptor = message_descriptor
|
||||
|
||||
def _SetFrozen(self) -> None:
|
||||
super()._SetFrozen()
|
||||
for val in self._values:
|
||||
val._SetFrozen()
|
||||
|
||||
def add(self, **kwargs: Any) -> _T:
|
||||
"""Adds a new element at the end of the list and returns it.
|
||||
|
||||
Keyword arguments may be used to initialize the element.
|
||||
"""
|
||||
self._AssureWritable()
|
||||
new_element = self._message_descriptor._concrete_class(**kwargs)
|
||||
new_element._SetListener(self._message_listener)
|
||||
self._values.append(new_element)
|
||||
|
|
@ -287,6 +312,7 @@ class RepeatedCompositeFieldContainer(BaseContainer[_T], MutableSequence[_T]):
|
|||
|
||||
def append(self, value: _T) -> None:
|
||||
"""Appends one element by copying the message."""
|
||||
self._AssureWritable()
|
||||
new_element = self._message_descriptor._concrete_class()
|
||||
new_element._SetListener(self._message_listener)
|
||||
new_element.CopyFrom(value)
|
||||
|
|
@ -296,6 +322,7 @@ class RepeatedCompositeFieldContainer(BaseContainer[_T], MutableSequence[_T]):
|
|||
|
||||
def insert(self, key: int, value: _T) -> None:
|
||||
"""Inserts the item at the specified position by copying."""
|
||||
self._AssureWritable()
|
||||
new_element = self._message_descriptor._concrete_class()
|
||||
new_element._SetListener(self._message_listener)
|
||||
new_element.CopyFrom(value)
|
||||
|
|
@ -308,6 +335,7 @@ class RepeatedCompositeFieldContainer(BaseContainer[_T], MutableSequence[_T]):
|
|||
|
||||
as this one, copying each individual message.
|
||||
"""
|
||||
self._AssureWritable()
|
||||
message_class = self._message_descriptor._concrete_class
|
||||
listener = self._message_listener
|
||||
values = self._values
|
||||
|
|
@ -326,15 +354,18 @@ class RepeatedCompositeFieldContainer(BaseContainer[_T], MutableSequence[_T]):
|
|||
|
||||
one, copying each individual message.
|
||||
"""
|
||||
self._AssureWritable()
|
||||
self.extend(other)
|
||||
|
||||
def remove(self, elem: _T) -> None:
|
||||
"""Removes an item from the list. Similar to list.remove()."""
|
||||
self._AssureWritable()
|
||||
self._values.remove(elem)
|
||||
self._message_listener.Modified()
|
||||
|
||||
def pop(self, key: Optional[int] = -1) -> _T:
|
||||
"""Removes and returns an item at a given index. Similar to list.pop()."""
|
||||
self._AssureWritable()
|
||||
value = self._values[key]
|
||||
self.__delitem__(key)
|
||||
return value
|
||||
|
|
@ -357,6 +388,7 @@ class RepeatedCompositeFieldContainer(BaseContainer[_T], MutableSequence[_T]):
|
|||
|
||||
def __delitem__(self, key: Union[int, slice]) -> None:
|
||||
"""Deletes the item at the specified position."""
|
||||
self._AssureWritable()
|
||||
del self._values[key]
|
||||
self._message_listener.Modified()
|
||||
|
||||
|
|
@ -382,6 +414,7 @@ class ScalarMap(MutableMapping[_K, _V]):
|
|||
'_values',
|
||||
'_message_listener',
|
||||
'_entry_descriptor',
|
||||
'_frozen',
|
||||
]
|
||||
|
||||
def __init__(
|
||||
|
|
@ -407,11 +440,21 @@ class ScalarMap(MutableMapping[_K, _V]):
|
|||
self._value_checker = value_checker
|
||||
self._entry_descriptor = entry_descriptor
|
||||
self._values = {}
|
||||
self._frozen = False
|
||||
|
||||
def _SetFrozen(self) -> None:
|
||||
self._frozen = True
|
||||
|
||||
def _AssureWritable(self) -> 'ScalarMap[_K, _V]':
|
||||
if self._frozen:
|
||||
raise TypeError('Map is frozen')
|
||||
return self
|
||||
|
||||
def __getitem__(self, key: _K) -> _V:
|
||||
try:
|
||||
return self._values[key]
|
||||
except KeyError:
|
||||
self._AssureWritable()
|
||||
key = self._key_checker.CheckValue(key)
|
||||
val = self._value_checker.DefaultValue()
|
||||
self._values[key] = val
|
||||
|
|
@ -441,12 +484,14 @@ class ScalarMap(MutableMapping[_K, _V]):
|
|||
return default
|
||||
|
||||
def __setitem__(self, key: _K, value: _V) -> _T:
|
||||
self._AssureWritable()
|
||||
checked_key = self._key_checker.CheckValue(key)
|
||||
checked_value = self._value_checker.CheckValue(value)
|
||||
self._values[checked_key] = checked_value
|
||||
self._message_listener.Modified()
|
||||
|
||||
def __delitem__(self, key: _K) -> None:
|
||||
self._AssureWritable()
|
||||
del self._values[key]
|
||||
self._message_listener.Modified()
|
||||
|
||||
|
|
@ -460,6 +505,7 @@ class ScalarMap(MutableMapping[_K, _V]):
|
|||
return repr(self._values)
|
||||
|
||||
def setdefault(self, key: _K, value: Optional[_V] = None) -> _V:
|
||||
self._AssureWritable()
|
||||
if value == None:
|
||||
raise ValueError('The value for scalar map setdefault must be set.')
|
||||
if key not in self._values:
|
||||
|
|
@ -467,6 +513,7 @@ class ScalarMap(MutableMapping[_K, _V]):
|
|||
return self[key]
|
||||
|
||||
def MergeFrom(self, other: 'ScalarMap[_K, _V]') -> None:
|
||||
self._AssureWritable()
|
||||
self._values.update(other._values)
|
||||
self._message_listener.Modified()
|
||||
|
||||
|
|
@ -479,6 +526,7 @@ class ScalarMap(MutableMapping[_K, _V]):
|
|||
|
||||
# This is defined in the abstract base, but we can do it much more cheaply.
|
||||
def clear(self) -> None:
|
||||
self._AssureWritable()
|
||||
self._values.clear()
|
||||
self._message_listener.Modified()
|
||||
|
||||
|
|
@ -496,6 +544,7 @@ class MessageMap(MutableMapping[_K, _V]):
|
|||
'_message_listener',
|
||||
'_message_descriptor',
|
||||
'_entry_descriptor',
|
||||
'_frozen',
|
||||
]
|
||||
|
||||
def __init__(
|
||||
|
|
@ -521,12 +570,24 @@ class MessageMap(MutableMapping[_K, _V]):
|
|||
self._key_checker = key_checker
|
||||
self._entry_descriptor = entry_descriptor
|
||||
self._values = {}
|
||||
self._frozen = False
|
||||
|
||||
def _SetFrozen(self) -> None:
|
||||
self._frozen = True
|
||||
for val in self._values.values():
|
||||
val._SetFrozen()
|
||||
|
||||
def _AssureWritable(self) -> 'MessageMap[_K, _V]':
|
||||
if self._frozen:
|
||||
raise TypeError('Map is immutable')
|
||||
return self
|
||||
|
||||
def __getitem__(self, key: _K) -> _V:
|
||||
key = self._key_checker.CheckValue(key)
|
||||
try:
|
||||
return self._values[key]
|
||||
except KeyError:
|
||||
self._AssureWritable()
|
||||
new_element = self._message_descriptor._concrete_class()
|
||||
new_element._SetListener(self._message_listener)
|
||||
self._values[key] = new_element
|
||||
|
|
@ -569,9 +630,11 @@ class MessageMap(MutableMapping[_K, _V]):
|
|||
return item in self._values
|
||||
|
||||
def __setitem__(self, key: _K, value: _V) -> NoReturn:
|
||||
self._AssureWritable()
|
||||
raise ValueError('May not set values directly, call my_map[key].foo = 5')
|
||||
|
||||
def __delitem__(self, key: _K) -> None:
|
||||
self._AssureWritable()
|
||||
key = self._key_checker.CheckValue(key)
|
||||
del self._values[key]
|
||||
self._message_listener.Modified()
|
||||
|
|
@ -586,12 +649,14 @@ class MessageMap(MutableMapping[_K, _V]):
|
|||
return repr(self._values)
|
||||
|
||||
def setdefault(self, key: _K, value: Optional[_V] = None) -> _V:
|
||||
self._AssureWritable()
|
||||
raise NotImplementedError(
|
||||
'Set message map value directly is not supported, call'
|
||||
' my_map[key].foo = 5'
|
||||
)
|
||||
|
||||
def MergeFrom(self, other: 'MessageMap[_K, _V]') -> None:
|
||||
self._AssureWritable()
|
||||
# pylint: disable=protected-access
|
||||
for key in other._values:
|
||||
# According to documentation: "When parsing from the wire or when merging,
|
||||
|
|
@ -611,6 +676,7 @@ class MessageMap(MutableMapping[_K, _V]):
|
|||
|
||||
# This is defined in the abstract base, but we can do it much more cheaply.
|
||||
def clear(self) -> None:
|
||||
self._AssureWritable()
|
||||
self._values.clear()
|
||||
self._message_listener.Modified()
|
||||
|
||||
|
|
|
|||
|
|
@ -280,14 +280,15 @@ class DescriptorTest(unittest.TestCase):
|
|||
self.my_service.GetOptions(), descriptor_pb2.ServiceOptions()
|
||||
)
|
||||
|
||||
@unittest.skipIf(
|
||||
api_implementation.Type() == 'python', 'Not fixed yet in pure Python'
|
||||
)
|
||||
@unittest.skipIf(api_implementation.Type() == 'cpp', 'Not fixed yet in C++')
|
||||
@unittest.skipIf(
|
||||
api_implementation.Type() == 'upb',
|
||||
'Needs to wait for a breaking change release in OSS'
|
||||
)
|
||||
@unittest.skipIf(
|
||||
api_implementation.Type() == 'python',
|
||||
'Needs to wait for a breaking change release in OSS'
|
||||
)
|
||||
def testModifyFrozenMessage(self):
|
||||
# At least upb raises TypeError Other 2 implementations will likely be
|
||||
# fixed to be consistent with upb.
|
||||
|
|
@ -351,9 +352,17 @@ class DescriptorTest(unittest.TestCase):
|
|||
# Extension dict mutation
|
||||
with self.assertRaises(immutability_error):
|
||||
message_options.Extensions[complex_opt1] = descriptor_pb2.MessageOptions()
|
||||
|
||||
message_opt1 = unittest_custom_options_pb2.message_opt1
|
||||
with self.assertRaises(immutability_error):
|
||||
message_options.Extensions[message_opt1] = -56
|
||||
|
||||
with self.assertRaises(immutability_error):
|
||||
del message_options.Extensions[complex_opt1]
|
||||
|
||||
with self.assertRaises(immutability_error):
|
||||
message_options.ClearExtension(complex_opt1)
|
||||
|
||||
# Map field mutations
|
||||
map_field = stub_submsg.my_map
|
||||
with self.assertRaises(immutability_error):
|
||||
|
|
|
|||
|
|
@ -93,6 +93,8 @@ class _ExtensionDict(object):
|
|||
# WARNING: We are relying on setdefault() being atomic. This is true
|
||||
# in CPython but we haven't investigated others. This warning appears
|
||||
# in several other locations in this file.
|
||||
if self._extended_message._frozen:
|
||||
result._SetFrozen()
|
||||
result = self._extended_message._fields.setdefault(extension_handle, result)
|
||||
|
||||
return result
|
||||
|
|
@ -134,6 +136,8 @@ class _ExtensionDict(object):
|
|||
|
||||
_VerifyExtensionHandle(self._extended_message, extension_handle)
|
||||
|
||||
self._extended_message._AssureWritable()
|
||||
|
||||
if (
|
||||
extension_handle.is_repeated
|
||||
or extension_handle.cpp_type == FieldDescriptor.CPPTYPE_MESSAGE
|
||||
|
|
|
|||
|
|
@ -254,6 +254,7 @@ def _AddSlots(message_descriptor, dictionary):
|
|||
'_listener_for_children',
|
||||
'__weakref__',
|
||||
'_oneofs',
|
||||
'_frozen',
|
||||
]
|
||||
|
||||
|
||||
|
|
@ -611,6 +612,7 @@ def _AddInitMethod(message_descriptor, cls):
|
|||
self._is_present_in_parent = False
|
||||
self._listener = message_listener_mod.NullMessageListener()
|
||||
self._listener_for_children = _Listener(self)
|
||||
self._frozen = False
|
||||
for field_name, field_value in kwargs.items():
|
||||
field = _GetFieldByName(message_descriptor, field_name)
|
||||
if field is None:
|
||||
|
|
@ -763,6 +765,8 @@ def _AddPropertiesForRepeatedField(field, cls):
|
|||
if field_value is None:
|
||||
# Construct a new object to represent this field.
|
||||
field_value = field._default_constructor(self)
|
||||
if self._frozen:
|
||||
field_value._SetFrozen()
|
||||
|
||||
# Atomically check if another thread has preempted us and, if not, swap
|
||||
# in the new object we just created. If someone has preempted us, we
|
||||
|
|
@ -814,6 +818,7 @@ def _AddPropertiesForNonRepeatedScalarField(field, cls):
|
|||
getter.__doc__ = 'Getter for %s.' % proto_field_name
|
||||
|
||||
def field_setter(self, new_value):
|
||||
self._AssureWritable()
|
||||
# pylint: disable=protected-access
|
||||
# Testing the value for truthiness captures all of the implicit presence
|
||||
# defaults (0, 0.0, enum 0, and False), except for -0.0.
|
||||
|
|
@ -871,6 +876,8 @@ def _AddPropertiesForNonRepeatedCompositeField(field, cls):
|
|||
if field_value is None:
|
||||
# Construct a new object to represent this field.
|
||||
field_value = field._default_constructor(self)
|
||||
if self._frozen:
|
||||
field_value._SetFrozen()
|
||||
|
||||
# Atomically check if another thread has preempted us and, if not, swap
|
||||
# in the new object we just created. If someone has preempted us, we
|
||||
|
|
@ -887,6 +894,7 @@ def _AddPropertiesForNonRepeatedCompositeField(field, cls):
|
|||
# We define a setter just so we can throw an exception with a more
|
||||
# helpful error message.
|
||||
def setter(self, new_value):
|
||||
self._AssureWritable()
|
||||
if field.message_type.full_name == 'google.protobuf.Timestamp':
|
||||
getter(self)
|
||||
self._fields[field].FromDatetime(new_value)
|
||||
|
|
@ -1013,6 +1021,7 @@ def _AddClearFieldMethod(message_descriptor, cls):
|
|||
"""Helper for _AddMessageMethods()."""
|
||||
|
||||
def ClearField(self, field_name):
|
||||
self._AssureWritable()
|
||||
try:
|
||||
field = message_descriptor.fields_by_name[field_name]
|
||||
except KeyError:
|
||||
|
|
@ -1054,6 +1063,7 @@ def _AddClearExtensionMethod(cls):
|
|||
"""Helper for _AddMessageMethods()."""
|
||||
|
||||
def ClearExtension(self, field_descriptor):
|
||||
self._AssureWritable()
|
||||
extension_dict._VerifyExtensionHandle(self, field_descriptor)
|
||||
|
||||
# Similar to ClearField(), above.
|
||||
|
|
@ -1328,6 +1338,7 @@ def _AddMergeFromStringMethod(message_descriptor, cls):
|
|||
"""Helper for _AddMessageMethods()."""
|
||||
|
||||
def MergeFromString(self, serialized):
|
||||
self._AssureWritable()
|
||||
serialized = memoryview(serialized)
|
||||
length = len(serialized)
|
||||
try:
|
||||
|
|
@ -1513,6 +1524,7 @@ def _AddMergeFromMethod(cls):
|
|||
CPPTYPE_MESSAGE = _FieldDescriptor.CPPTYPE_MESSAGE
|
||||
|
||||
def MergeFrom(self, msg):
|
||||
self._AssureWritable()
|
||||
if not isinstance(msg, cls):
|
||||
raise TypeError(
|
||||
'Parameter to MergeFrom() must be instance of same class: '
|
||||
|
|
@ -1576,6 +1588,7 @@ def _AddWhichOneofMethod(message_descriptor, cls):
|
|||
|
||||
|
||||
def _Clear(self):
|
||||
self._AssureWritable()
|
||||
# Clear fields.
|
||||
self._fields = {}
|
||||
self._unknown_fields = ()
|
||||
|
|
@ -1584,6 +1597,19 @@ def _Clear(self):
|
|||
self._Modified()
|
||||
|
||||
|
||||
def _SetFrozen(self):
|
||||
self._frozen = True
|
||||
for value in self._fields.values():
|
||||
if hasattr(value, '_SetFrozen'):
|
||||
value._SetFrozen()
|
||||
|
||||
|
||||
def _AssureWritable(self):
|
||||
if self._frozen:
|
||||
raise TypeError('Message is immutable')
|
||||
return self
|
||||
|
||||
|
||||
def _UnknownFields(self):
|
||||
raise NotImplementedError(
|
||||
'Please use the add-on feaure '
|
||||
|
|
@ -1593,6 +1619,7 @@ def _UnknownFields(self):
|
|||
|
||||
|
||||
def _DiscardUnknownFields(self):
|
||||
self._AssureWritable()
|
||||
self._unknown_fields = []
|
||||
for field, value in self.ListFields():
|
||||
if field.cpp_type == _FieldDescriptor.CPPTYPE_MESSAGE:
|
||||
|
|
@ -1638,6 +1665,8 @@ def _AddMessageMethods(message_descriptor, cls):
|
|||
cls.Clear = _Clear
|
||||
cls.DiscardUnknownFields = _DiscardUnknownFields
|
||||
cls._SetListener = _SetListener
|
||||
cls._SetFrozen = _SetFrozen
|
||||
cls._AssureWritable = _AssureWritable
|
||||
|
||||
|
||||
def _AddPrivateHelperMethods(message_descriptor, cls):
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue