dep-protobuf/python/google/protobuf/internal/thread_safe_test.py
Joshua Haberman 6cbc7593bf Fixed race in GetMessageClass/RegisterMessageClass under free threading.
The race condition was resolved by adding a mutex lock to all accesses of the cache.

PiperOrigin-RevId: 912595698
2026-05-08 10:47:54 -07:00

402 lines
11 KiB
Python

# Protocol Buffers - Google's data interchange format
# Copyright 2008 Google Inc. All rights reserved.
#
# Use of this source code is governed by a BSD-style
# license that can be found in the LICENSE file or at
# https://developers.google.com/open-source/licenses/bsd
"""Unittest for thread safe"""
import threading
import time
import timeit
import unittest
from google.protobuf import descriptor_pb2
from google.protobuf import descriptor_pool
from google.protobuf import message_factory
from google.protobuf.internal import api_implementation
from google.protobuf.internal import testing_refleaks
from google.protobuf import unittest_pb2
from google.protobuf import unittest_proto3_pb2
# Enable this to run the benchmarks.
ALSO_RUN_BENCHMARKS = False
@testing_refleaks.TestCase
class ThreadSafeTest(unittest.TestCase):
def setUp(self):
self.success = 0
def testFieldDecodersDataRace(self):
msg = unittest_pb2.TestAllTypes(optional_int32=1)
serialized_data = msg.SerializeToString()
lock = threading.Lock()
def ParseMessage():
parsed_msg = unittest_pb2.TestAllTypes()
time.sleep(0.005)
parsed_msg.ParseFromString(serialized_data)
with lock:
if msg == parsed_msg:
self.success += 1
field_des = unittest_pb2.TestAllTypes.DESCRIPTOR.fields_by_name[
'optional_int32'
]
count = 1000
for x in range(0, count):
# delete the _decoders because only the first time parse the field
# may cause data race.
if hasattr(field_des, '_decoders'):
delattr(field_des, '_decoders')
thread1 = threading.Thread(target=ParseMessage)
thread2 = threading.Thread(target=ParseMessage)
thread1.start()
thread2.start()
thread1.join()
thread2.join()
self.assertEqual(count * 2, self.success)
# This caused a Dealloc()/Dealloc() race.
@unittest.skipIf(
api_implementation.Type() == 'upb',
'Upb has not been fixed to handle this case.',
)
def testGetType(self):
def GetType():
msg = unittest_proto3_pb2.TestAllTypes(
optional_nested_message=unittest_proto3_pb2.TestAllTypes.NestedMessage(
bb=1000
),
optional_nested_enum=unittest_proto3_pb2.TestAllTypes.NestedEnum.ZERO,
)
msges = [msg] * 100
for m in msges:
# Fails in this line:
unittest_proto3_pb2.TestAllTypes.NestedEnum.Name(m.optional_nested_enum)
threads = []
for i in range(100):
thread = threading.Thread(target=GetType)
threads.append(thread)
thread.start()
for thread in threads:
thread.join()
# This caused a race between constructing and using the type.
@unittest.skipIf(
api_implementation.Type() == 'upb',
'Upb has not been fixed to handle this case.',
)
def testInitType(self):
def InitType():
array = []
for i in range(100):
array.append(
unittest_proto3_pb2.TestAllTypes(
optional_nested_message=unittest_proto3_pb2.TestAllTypes.NestedMessage(
bb=1000
),
optional_nested_enum=unittest_proto3_pb2.TestAllTypes.NestedEnum.FOO,
)
)
threads = []
for i in range(100):
thread = threading.Thread(target=InitType)
threads.append(thread)
thread.start()
for thread in threads:
thread.join()
@unittest.skipIf(
api_implementation.Type() == 'upb',
'Upb has not been fixed to handle this case.',
)
def testConcurrentSubMessageAccess(self):
msg = unittest_proto3_pb2.TestAllTypes(
optional_nested_message=unittest_proto3_pb2.TestAllTypes.NestedMessage(
bb=1000
)
)
def AccessSubMessage():
for _ in range(100):
_ = msg.optional_nested_message.bb
threads = []
for i in range(100):
thread = threading.Thread(target=AccessSubMessage)
threads.append(thread)
thread.start()
for thread in threads:
thread.join()
@unittest.skipIf(
api_implementation.Type() == 'upb',
'Upb has not been fixed to handle this case.',
)
def testConcurrentRepeatedMessageAccess(self):
variable = unittest_proto3_pb2.TestAllTypes()
def UseVariable():
for _ in range(1000):
_ = variable.repeated_nested_message
threads = []
for i in range(100):
thread = threading.Thread(target=UseVariable)
threads.append(thread)
thread.start()
for thread in threads:
thread.join()
@unittest.skipIf(
api_implementation.Type() == 'upb',
'Upb has not been fixed to handle this case.',
)
def testConcurrentRepeatedPrimitiveAccess(self):
variable = unittest_proto3_pb2.TestAllTypes()
variable.repeated_float.append(1.0)
def UseVariable():
for _ in range(1000):
_ = variable.repeated_float
threads = []
for i in range(100):
thread = threading.Thread(target=UseVariable)
threads.append(thread)
thread.start()
for thread in threads:
thread.join()
@unittest.skipIf(
api_implementation.Type() == 'upb',
'Upb has not been fixed to handle this case.',
)
def testConcurrentSingularFieldAccess(self):
variable = unittest_proto3_pb2.TestAllTypes()
def UseVariable():
for _ in range(1000):
_ = variable.optional_int32
_ = variable.optional_string
threads = []
for i in range(100):
thread = threading.Thread(target=UseVariable)
threads.append(thread)
thread.start()
for thread in threads:
thread.join()
@unittest.skipIf(
api_implementation.Type() == 'upb',
'Upb has not been fixed to handle this case.',
)
def testConcurrentRepeatedMessageAccess2(self):
msg = unittest_proto3_pb2.TestAllTypes(
repeated_nested_message=[
unittest_proto3_pb2.TestAllTypes.NestedMessage(bb=1)
]
)
def UseVariable():
for _ in range(1000):
for nested in msg.repeated_nested_message:
pass
threads = []
for _ in range(100):
thread = threading.Thread(target=UseVariable)
threads.append(thread)
thread.start()
for thread in threads:
thread.join()
class FreeThreadingTest(unittest.TestCase):
def RunThreads(self, thread_size, func):
threads = []
for i in range(0, thread_size):
threads.append(threading.Thread(target=func))
for thread in threads:
thread.start()
for thread in threads:
thread.join()
def testDoNothing(self):
thread_size = 10
def DoNothing():
return
self.RunThreads(thread_size, DoNothing)
def testDescriptorPoolMap(self):
thread_size = 20
self.success_count = 0
lock = threading.Lock()
def CreatePool():
def DoCreate():
pool = descriptor_pool.DescriptorPool()
file_proto = descriptor_pb2.FileDescriptorProto(name='foo')
message_proto = file_proto.message_type.add(name='SomeMessage')
message_proto.field.add(
name='int_field',
number=1,
type=descriptor_pb2.FieldDescriptorProto.TYPE_INT32,
label=descriptor_pb2.FieldDescriptorProto.LABEL_OPTIONAL,
)
pool.Add(file_proto)
desc = pool.FindMessageTypeByName('SomeMessage')
msg = message_factory.GetMessageClass(desc)()
msg.int_field = 1
DoCreate()
with lock:
self.success_count += 1
self.RunThreads(thread_size, CreatePool)
self.assertEqual(thread_size, self.success_count)
@unittest.skipIf(
api_implementation.Type() == 'upb',
'Upb has not been fixed to handle this case.',
)
def testConcurrentGetFieldValueRace(self):
"""Reproduces a data race in GetFieldValue due to lazy initialization."""
def AccessFields(msg, barrier) -> None:
barrier.wait()
# This access triggers GetFieldValue and lazy initialization
# of the composite_fields map in CMessage.
_ = msg.optional_nested_message
for _ in range(100):
threads = []
msg = unittest_proto3_pb2.TestAllTypes()
# Use a barrier to ensure all threads hit the GetFieldValue call
# at nearly the same time, maximizing the race window.
barrier = threading.Barrier(10)
for _ in range(10):
thread = threading.Thread(target=AccessFields, args=(msg, barrier))
threads.append(thread)
thread.start()
for thread in threads:
thread.join()
@unittest.skipIf(
api_implementation.Type() == 'upb',
'Upb has not been fixed to handle this case.',
)
def testConcurrentGetOptionsRace(self):
"""Reproduces a data race in GetOptions."""
def AccessOptions(barrier):
barrier.wait()
_ = unittest_proto3_pb2.TestAllTypes.DESCRIPTOR.GetOptions()
for _ in range(100):
threads = []
barrier = threading.Barrier(20)
for _ in range(20):
thread = threading.Thread(target=AccessOptions, args=(barrier,))
threads.append(thread)
thread.start()
for thread in threads:
thread.join()
@unittest.skipIf(
api_implementation.Type() == 'upb',
'Upb has not been fixed to handle this case.',
)
def testConcurrentGetAndRegisterMessageClassDataRace(self):
"""Reproduces the data race in GetMessageClass/RegisterMessageClass."""
pool = descriptor_pool.DescriptorPool()
# Create and register a base Message class to look up
file_proto = descriptor_pb2.FileDescriptorProto(name='base.proto')
file_proto.message_type.add(name='BaseMessage')
pool.Add(file_proto)
base_desc = pool.FindMessageTypeByName('BaseMessage')
message_factory.GetMessageClass(base_desc)
# Pre-create descriptors for modification
num_descriptors = 500
descriptors = []
for i in range(num_descriptors):
name = f'DynamicMessage_{i}'
f_proto = descriptor_pb2.FileDescriptorProto(name=f'{name}.proto')
f_proto.message_type.add(name=name)
pool.Add(f_proto)
descriptors.append(pool.FindMessageTypeByName(name))
barrier = threading.Barrier(10)
def Task(thread_id: int):
barrier.wait()
if thread_id % 2 == 0:
# Reader thread: repeatedly looks up the existing message class
for _ in range(200):
message_factory.GetMessageClass(base_desc)
else:
# Writer thread: registers new message classes concurrently
start_idx = (thread_id // 2) * 100
for i in range(start_idx, start_idx + 100):
message_factory.GetMessageClass(descriptors[i])
threads = []
for i in range(10):
threads.append(threading.Thread(target=Task, args=(i,)))
for thread in threads:
thread.start()
for thread in threads:
thread.join()
@unittest.skipIf(not ALSO_RUN_BENCHMARKS, 'Benchmarks are disabled.')
def testConcurrentGetOptionsBenchmark(self):
"""Benchmarks concurrent GetOptions calls."""
if ALSO_RUN_BENCHMARKS:
def AccessOptions():
for _ in range(1000000):
_ = unittest_proto3_pb2.TestAllTypes.DESCRIPTOR.GetOptions()
def RunAllThreads():
self.RunThreads(20, AccessOptions)
duration = timeit.timeit(RunAllThreads, number=10)
duration_ms = duration * 1000
print(
'ConcurrentGetOptionsBenchmark (20 threads x 1000000 calls x 10'
f' runs): {duration_ms:.2f}ms'
)
else:
print('Skipping benchmark in non-benchmark mode.')
if __name__ == '__main__':
unittest.main()