[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:
Runze Wang 2026-06-03 10:16:14 -07:00 committed by Copybara-Service
parent 791fbe2249
commit f5738bb5f0
4 changed files with 112 additions and 4 deletions

View file

@ -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()

View file

@ -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):

View file

@ -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

View file

@ -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):