From f5738bb5f0a27580f50c86ea2e86102b89e214ac Mon Sep 17 00:00:00 2001 From: Runze Wang Date: Wed, 3 Jun 2026 10:16:14 -0700 Subject: [PATCH] [py/pure python] Implement Make GetOptions() return immutable options. Raise a TypeError when options returned GetOptions() by is mutated. PiperOrigin-RevId: 926114776 --- python/google/protobuf/internal/containers.py | 68 ++++++++++++++++++- .../protobuf/internal/descriptor_test.py | 15 +++- .../protobuf/internal/extension_dict.py | 4 ++ .../protobuf/internal/python_message.py | 29 ++++++++ 4 files changed, 112 insertions(+), 4 deletions(-) diff --git a/python/google/protobuf/internal/containers.py b/python/google/protobuf/internal/containers.py index f06a35d5e7..76594caaaf 100755 --- a/python/google/protobuf/internal/containers.py +++ b/python/google/protobuf/internal/containers.py @@ -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() diff --git a/python/google/protobuf/internal/descriptor_test.py b/python/google/protobuf/internal/descriptor_test.py index 325ffa9c5c..3e715ba67e 100755 --- a/python/google/protobuf/internal/descriptor_test.py +++ b/python/google/protobuf/internal/descriptor_test.py @@ -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): diff --git a/python/google/protobuf/internal/extension_dict.py b/python/google/protobuf/internal/extension_dict.py index ad246b61f3..5e6a132358 100644 --- a/python/google/protobuf/internal/extension_dict.py +++ b/python/google/protobuf/internal/extension_dict.py @@ -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 diff --git a/python/google/protobuf/internal/python_message.py b/python/google/protobuf/internal/python_message.py index 49c6f55701..c86515d8c3 100755 --- a/python/google/protobuf/internal/python_message.py +++ b/python/google/protobuf/internal/python_message.py @@ -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):