mirror of
https://github.com/protocolbuffers/protobuf
synced 2026-08-26 02:23:14 -04:00
Internal change.
PiperOrigin-RevId: 941222035
This commit is contained in:
parent
ee6026fbbd
commit
dc4d738909
2 changed files with 174 additions and 7 deletions
|
|
@ -531,16 +531,85 @@ class MessageTest(unittest.TestCase):
|
|||
self.assertEqual([1, 2, 3, 4], msg.payload.repeated_int32)
|
||||
|
||||
def testRepeatedFieldSelfSliceAssignment(self, message_module):
|
||||
msg = message_module.NestedTestAllTypes()
|
||||
msg.payload.repeated_int32[:] = [1, 2, 3, 4]
|
||||
msg.payload.repeated_int32[:] = msg.payload.repeated_int32
|
||||
msg = message_module.NestedTestAllTypes()
|
||||
for field_name in [
|
||||
'repeated_int32',
|
||||
'repeated_int64',
|
||||
'repeated_uint32',
|
||||
'repeated_uint64',
|
||||
'repeated_sint32',
|
||||
'repeated_sint64',
|
||||
'repeated_fixed32',
|
||||
'repeated_fixed64',
|
||||
'repeated_sfixed32',
|
||||
'repeated_sfixed64',
|
||||
]:
|
||||
field = getattr(msg.payload, field_name)
|
||||
field[:] = [1, 2, 3, 4]
|
||||
field[:] = field
|
||||
self.assertEqual([1, 2, 3, 4], field)
|
||||
field[:] = field[1:-1]
|
||||
self.assertEqual([2, 3], field)
|
||||
for field_name in [
|
||||
'repeated_float',
|
||||
'repeated_double',
|
||||
]:
|
||||
field = getattr(msg.payload, field_name)
|
||||
field[:] = [1.25, 2.25, 3.25, 4.25]
|
||||
field[:] = field
|
||||
self.assertEqual([1.25, 2.25, 3.25, 4.25], field)
|
||||
|
||||
def testRepeatedFieldExtendWithPartialSuccess(self, message_module):
|
||||
msg = message_module.NestedTestAllTypes()
|
||||
msg.payload.repeated_int32[:] = [1, 2, 3, 4]
|
||||
with self.assertRaises(ValueError):
|
||||
msg.payload.repeated_int32.extend([4, 5, 6, 2**34])
|
||||
if api_implementation.Type() == 'cpp':
|
||||
self.assertEqual([1, 2, 3, 4, 4, 5, 6], msg.payload.repeated_int32)
|
||||
else:
|
||||
self.assertEqual([1, 2, 3, 4], msg.payload.repeated_int32)
|
||||
|
||||
def testRepeatedFieldSubSliceAssignment(self, message_module):
|
||||
msg = message_module.NestedTestAllTypes()
|
||||
msg.payload.repeated_int32[:] = range(1, 6)
|
||||
msg.payload.repeated_int32[1:3] = msg.payload.repeated_int32[2:4]
|
||||
self.assertEqual([1, 3, 4, 4, 5], msg.payload.repeated_int32)
|
||||
msg.payload.repeated_int32.extend(msg.payload.repeated_int32[1:3])
|
||||
self.assertEqual([1, 3, 4, 4, 5, 3, 4], msg.payload.repeated_int32)
|
||||
|
||||
def testRepeatedFieldDifferentTypeSliceAssignment(self, message_module):
|
||||
msg1 = message_module.NestedTestAllTypes()
|
||||
msg2 = message_module.NestedTestAllTypes()
|
||||
# int64 -> int32
|
||||
msg2.payload.repeated_int64[:] = [1, 2, 3, 4]
|
||||
msg1.payload.repeated_int32[:] = msg2.payload.repeated_int64
|
||||
self.assertEqual([1, 2, 3, 4], msg1.payload.repeated_int32)
|
||||
# int32 -> int64
|
||||
msg2.payload.repeated_int32[:] = [1, 2, 3, 4]
|
||||
msg1.payload.repeated_int64[:] = msg2.payload.repeated_int32
|
||||
self.assertEqual([1, 2, 3, 4], msg1.payload.repeated_int64)
|
||||
# int64 overflow -> int32
|
||||
msg2.payload.repeated_int64[:] = [1, 2, 3, 2**35]
|
||||
with self.assertRaises((ValueError, OverflowError, TypeError)):
|
||||
msg1.payload.repeated_int32[:] = msg2.payload.repeated_int64
|
||||
# double -> float
|
||||
msg2.payload.repeated_double[:] = [1.5, 2.5, 3.5]
|
||||
msg1.payload.repeated_float[:] = msg2.payload.repeated_double
|
||||
self.assertEqual([1.5, 2.5, 3.5], msg1.payload.repeated_float)
|
||||
# float -> double
|
||||
msg2.payload.repeated_float[:] = [1.5, 2.5, 3.5]
|
||||
msg1.payload.repeated_double[:] = msg2.payload.repeated_float
|
||||
self.assertEqual([1.5, 2.5, 3.5], msg1.payload.repeated_double)
|
||||
|
||||
msg2.payload.repeated_double[:] = [1.5, 2.5, 1e300]
|
||||
msg1.payload.repeated_float[:] = msg2.payload.repeated_double
|
||||
self.assertEqual([1.5, 2.5, float('inf')], msg1.payload.repeated_float)
|
||||
|
||||
def testRepeatedFieldSelfExtend(self, message_module):
|
||||
msg = message_module.NestedTestAllTypes()
|
||||
msg.payload.repeated_int32[:] = [1, 2, 3, 4]
|
||||
msg.payload.repeated_int32.extend(msg.payload.repeated_int32)
|
||||
self.assertEqual([1, 2, 3, 4] * 2, msg.payload.repeated_int32)
|
||||
msg = message_module.NestedTestAllTypes()
|
||||
msg.payload.repeated_int32[:] = [1, 2, 3, 4]
|
||||
msg.payload.repeated_int32.extend(msg.payload.repeated_int32)
|
||||
self.assertEqual([1, 2, 3, 4] * 2, msg.payload.repeated_int32)
|
||||
|
||||
def testAssignOutOfRange(self, message_module):
|
||||
msg = message_module.NestedTestAllTypes()
|
||||
|
|
|
|||
|
|
@ -220,6 +220,32 @@ class NumpyIntProtoTest(unittest.TestCase):
|
|||
with self.assertRaises(TypeError):
|
||||
message.optional_int64 = np_22_float_array
|
||||
|
||||
def testRepeatedFieldSelfSliceAssignment(self):
|
||||
msg = unittest_pb2.NestedTestAllTypes()
|
||||
msg.payload.repeated_int32[:] = np.arange(4, dtype=np.int32)
|
||||
msg.payload.repeated_int32[:] = np.asarray(msg.payload.repeated_int32)
|
||||
self.assertEqual([0, 1, 2, 3], msg.payload.repeated_int32)
|
||||
|
||||
def testNumpyArrayIsMutableCopy(self):
|
||||
msg = unittest_pb2.NestedTestAllTypes()
|
||||
msg.payload.repeated_int32[:] = np.arange(4, dtype=np.int32)
|
||||
arr = np.asarray(msg.payload.repeated_int32)
|
||||
arr[0] = 100
|
||||
self.assertEqual([0, 1, 2, 3], msg.payload.repeated_int32)
|
||||
np.testing.assert_equal([100, 1, 2, 3], arr)
|
||||
|
||||
def testNumpyDifferentIntTypeSliceAssignment(self):
|
||||
msg = unittest_pb2.NestedTestAllTypes()
|
||||
# int64 -> int32
|
||||
msg.payload.repeated_int32[:] = np.arange(4, dtype=np.int64)
|
||||
self.assertEqual([0, 1, 2, 3], msg.payload.repeated_int32)
|
||||
# int32 -> int64
|
||||
msg.payload.repeated_int64[:] = np.arange(4, dtype=np.int32)
|
||||
self.assertEqual([0, 1, 2, 3], msg.payload.repeated_int64)
|
||||
# int64 overflow -> int32
|
||||
with self.assertRaises((ValueError, OverflowError, TypeError)):
|
||||
msg.payload.repeated_int32[:] = np.array([0, 1, 2, 2**35], dtype=np.int64)
|
||||
|
||||
|
||||
@testing_refleaks.TestCase
|
||||
class NumpyFloatProtoTest(unittest.TestCase):
|
||||
|
|
@ -274,6 +300,20 @@ class NumpyFloatProtoTest(unittest.TestCase):
|
|||
with self.assertRaises(TypeError):
|
||||
message.optional_float = np_22_object_array_float
|
||||
|
||||
def testNumpyDifferentFloatTypeSliceAssignment(self):
|
||||
msg = unittest_pb2.NestedTestAllTypes()
|
||||
# float64 -> float32
|
||||
msg.payload.repeated_float[:] = np.array([1.5, 2.5, 3.5], dtype=np.float64)
|
||||
self.assertEqual([1.5, 2.5, 3.5], msg.payload.repeated_float)
|
||||
# float32 -> float64
|
||||
msg.payload.repeated_double[:] = np.array([1.5, 2.5, 3.5], dtype=np.float32)
|
||||
self.assertEqual([1.5, 2.5, 3.5], msg.payload.repeated_double)
|
||||
# float64 overflow -> float32
|
||||
msg.payload.repeated_float[:] = np.array(
|
||||
[1.5, 2.5, 1e300], dtype=np.float64
|
||||
)
|
||||
self.assertEqual([1.5, 2.5, float('inf')], msg.payload.repeated_float)
|
||||
|
||||
|
||||
@testing_refleaks.TestCase
|
||||
class NumpyBoolProtoTest(unittest.TestCase):
|
||||
|
|
@ -749,6 +789,64 @@ class NumpyBindingTest(parameterized.TestCase):
|
|||
arr = np.array(m.repeated_int32, order='F')
|
||||
np.testing.assert_equal(arr, np.array([1, 2, 3]))
|
||||
|
||||
@parameterized.product(
|
||||
message_module=[unittest_pb2, unittest_proto3_arena_pb2],
|
||||
field_name=[
|
||||
'repeated_int32',
|
||||
'repeated_int64',
|
||||
'repeated_uint32',
|
||||
'repeated_uint64',
|
||||
'repeated_sint32',
|
||||
'repeated_sint64',
|
||||
'repeated_fixed32',
|
||||
'repeated_fixed64',
|
||||
'repeated_sfixed32',
|
||||
'repeated_sfixed64',
|
||||
],
|
||||
dtype=[
|
||||
np.int8,
|
||||
np.int16,
|
||||
np.int32,
|
||||
np.int64,
|
||||
np.uint8,
|
||||
np.uint16,
|
||||
np.uint32,
|
||||
np.uint64,
|
||||
],
|
||||
)
|
||||
def test_assign_integer_numpy_array_to_repeated(
|
||||
self, message_module, field_name, dtype
|
||||
):
|
||||
m = message_module.TestAllTypes()
|
||||
field = getattr(m, field_name)
|
||||
arr = np.array([0, 1, 2, 3], dtype=dtype)
|
||||
field[:] = arr
|
||||
self.assertEqual([0, 1, 2, 3], field)
|
||||
field[1:-1] = arr
|
||||
self.assertEqual([0, 0, 1, 2, 3, 3], field)
|
||||
|
||||
@parameterized.product(
|
||||
message_module=[unittest_pb2, unittest_proto3_arena_pb2],
|
||||
field_name=[
|
||||
'repeated_float',
|
||||
'repeated_double',
|
||||
],
|
||||
dtype=[
|
||||
np.float32,
|
||||
np.float64,
|
||||
],
|
||||
)
|
||||
def test_assign_float_numpy_array_to_repeated(
|
||||
self, message_module, field_name, dtype
|
||||
):
|
||||
m = message_module.TestAllTypes()
|
||||
field = getattr(m, field_name)
|
||||
arr = np.array([1.5, 2.5, 3.5], dtype=dtype)
|
||||
field[:] = arr
|
||||
self.assertEqual([1.5, 2.5, 3.5], field)
|
||||
field[1:-1] = arr
|
||||
self.assertEqual([1.5, 1.5, 2.5, 3.5, 3.5], field)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue