diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/streamed.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/streamed.py index d16955d88abb..c47cc0ef0a17 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/streamed.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/_async/streamed.py @@ -193,9 +193,9 @@ async def _consume_next(self): @CrossSync.convert(sync_name="__iter__") async def __aiter__(self): while True: - iter_rows, self._rows[:] = self._rows[:], () - while iter_rows: - yield iter_rows.pop(0) + iter_rows, self._rows = self._rows, [] + for row in iter_rows: + yield row if self._done: return try: diff --git a/packages/google-cloud-spanner/google/cloud/spanner_v1/streamed.py b/packages/google-cloud-spanner/google/cloud/spanner_v1/streamed.py index 8facd015151d..a92f008f5e32 100644 --- a/packages/google-cloud-spanner/google/cloud/spanner_v1/streamed.py +++ b/packages/google-cloud-spanner/google/cloud/spanner_v1/streamed.py @@ -168,9 +168,9 @@ def _consume_next(self): def __iter__(self): while True: - iter_rows, self._rows[:] = (self._rows[:], ()) - while iter_rows: - yield iter_rows.pop(0) + iter_rows, self._rows = (self._rows, []) + for row in iter_rows: + yield row if self._done: return try: diff --git a/packages/google-cloud-spanner/tests/unit/_async/test_streamed.py b/packages/google-cloud-spanner/tests/unit/_async/test_streamed.py index d8939cdba4a1..f3ec2bb4d0cb 100644 --- a/packages/google-cloud-spanner/tests/unit/_async/test_streamed.py +++ b/packages/google-cloud-spanner/tests/unit/_async/test_streamed.py @@ -1115,6 +1115,121 @@ async def test___iter___w_existing_rows_read(self): self.assertEqual(streamed._current_row, []) self.assertIsNone(streamed._pending_chunk) + @CrossSync.pytest + async def test___iter___large_batch(self): + from google.cloud.spanner_v1 import TypeCode + + fields = [ + self._make_scalar_field("id", TypeCode.INT64), + self._make_scalar_field("name", TypeCode.STRING), + ] + metadata = self._make_result_set_metadata(fields) + expected_rows = [[i, f"name_{i}"] for i in range(500)] + values = [self._make_value(cell) for row in expected_rows for cell in row] + + result_set = self._make_partial_result_set(values, metadata=metadata) + iterator = _MockCancellableIterator(result_set) + streamed = self._make_one(iterator) + found = [row async for row in streamed] + self.assertEqual(found, expected_rows) + self.assertEqual([row async for row in streamed], []) + + @CrossSync.pytest + async def test___iter___stepwise_consumption(self): + from google.cloud.spanner_v1 import TypeCode + + fields = [ + self._make_scalar_field("id", TypeCode.INT64), + self._make_scalar_field("name", TypeCode.STRING), + ] + metadata = self._make_result_set_metadata(fields) + expected_rows = [[i, f"name_{i}"] for i in range(20)] + values = [self._make_value(cell) for row in expected_rows for cell in row] + + result_set = self._make_partial_result_set(values, metadata=metadata) + iterator = _MockCancellableIterator(result_set) + streamed = self._make_one(iterator) + stream_iter = streamed.__aiter__() + first_five = [await stream_iter.__anext__() for _ in range(5)] + self.assertEqual(first_five, expected_rows[:5]) + remaining = [row async for row in stream_iter] + self.assertEqual(remaining, expected_rows[5:]) + + @CrossSync.pytest + async def test___iter___stepwise_across_chunks(self): + from google.cloud.spanner_v1 import TypeCode + + fields = [ + self._make_scalar_field("id", TypeCode.INT64), + self._make_scalar_field("name", TypeCode.STRING), + ] + metadata = self._make_result_set_metadata(fields) + chunk1_rows = [[i, f"name_{i}"] for i in range(10)] + chunk2_rows = [[i, f"name_{i}"] for i in range(10, 20)] + values1 = [self._make_value(cell) for row in chunk1_rows for cell in row] + values2 = [self._make_value(cell) for row in chunk2_rows for cell in row] + + result_set1 = self._make_partial_result_set(values1, metadata=metadata) + result_set2 = self._make_partial_result_set(values2) + iterator = _MockCancellableIterator(result_set1, result_set2) + streamed = self._make_one(iterator) + stream_iter = streamed.__aiter__() + first_part = [await stream_iter.__anext__() for _ in range(5)] + self.assertEqual(first_part, chunk1_rows[:5]) + middle_part = [await stream_iter.__anext__() for _ in range(10)] + self.assertEqual(middle_part, chunk1_rows[5:] + chunk2_rows[:5]) + final_part = [row async for row in stream_iter] + self.assertEqual(final_part, chunk2_rows[5:]) + + @CrossSync.pytest + async def test___iter___early_break(self): + from google.cloud.spanner_v1 import TypeCode + + fields = [ + self._make_scalar_field("id", TypeCode.INT64), + self._make_scalar_field("name", TypeCode.STRING), + ] + metadata = self._make_result_set_metadata(fields) + expected_rows = [[i, f"name_{i}"] for i in range(10)] + values = [self._make_value(cell) for row in expected_rows for cell in row] + + result_set = self._make_partial_result_set(values, metadata=metadata) + iterator = _MockCancellableIterator(result_set) + streamed = self._make_one(iterator) + consumed = [] + async for row in streamed: + consumed.append(row) + if len(consumed) == 3: + break + + self.assertEqual(consumed, expected_rows[:3]) + + @CrossSync.pytest + async def test___iter___mid_stream_error(self): + from google.cloud.spanner_v1 import TypeCode + + fields = [ + self._make_scalar_field("id", TypeCode.INT64), + self._make_scalar_field("name", TypeCode.STRING), + ] + metadata = self._make_result_set_metadata(fields) + chunk1_rows = [[i, f"name_{i}"] for i in range(5)] + values1 = [self._make_value(cell) for row in chunk1_rows for cell in row] + result_set1 = self._make_partial_result_set(values1, metadata=metadata) + + async def mock_iterator(): + yield result_set1 + raise RuntimeError("Stream error midway") + + streamed = self._make_one(mock_iterator()) + consumed = [] + with self.assertRaises(RuntimeError) as context: + async for row in streamed: + consumed.append(row) + + self.assertEqual(consumed, chunk1_rows) + self.assertIn("Stream error midway", str(context.exception)) + class _MockCancellableIterator(object): cancel_calls = 0 diff --git a/packages/google-cloud-spanner/tests/unit/test_streamed.py b/packages/google-cloud-spanner/tests/unit/test_streamed.py index 7cd505be5471..3d3ba709145d 100644 --- a/packages/google-cloud-spanner/tests/unit/test_streamed.py +++ b/packages/google-cloud-spanner/tests/unit/test_streamed.py @@ -1035,6 +1035,116 @@ def test___iter___w_existing_rows_read(self): self.assertEqual(streamed._current_row, []) self.assertIsNone(streamed._pending_chunk) + def test___iter___large_batch(self): + from google.cloud.spanner_v1 import TypeCode + + fields = [ + self._make_scalar_field("id", TypeCode.INT64), + self._make_scalar_field("name", TypeCode.STRING), + ] + metadata = self._make_result_set_metadata(fields) + expected_rows = [[i, f"name_{i}"] for i in range(500)] + values = [self._make_value(cell) for row in expected_rows for cell in row] + + result_set = self._make_partial_result_set(values, metadata=metadata) + iterator = _MockCancellableIterator(result_set) + streamed = self._make_one(iterator) + found = list(streamed) + self.assertEqual(found, expected_rows) + self.assertEqual(list(streamed), []) + + def test___iter___stepwise_consumption(self): + from google.cloud.spanner_v1 import TypeCode + + fields = [ + self._make_scalar_field("id", TypeCode.INT64), + self._make_scalar_field("name", TypeCode.STRING), + ] + metadata = self._make_result_set_metadata(fields) + expected_rows = [[i, f"name_{i}"] for i in range(20)] + values = [self._make_value(cell) for row in expected_rows for cell in row] + + result_set = self._make_partial_result_set(values, metadata=metadata) + iterator = _MockCancellableIterator(result_set) + streamed = self._make_one(iterator) + stream_iter = iter(streamed) + first_five = [next(stream_iter) for _ in range(5)] + self.assertEqual(first_five, expected_rows[:5]) + remaining = list(stream_iter) + self.assertEqual(remaining, expected_rows[5:]) + + def test___iter___stepwise_across_chunks(self): + from google.cloud.spanner_v1 import TypeCode + + fields = [ + self._make_scalar_field("id", TypeCode.INT64), + self._make_scalar_field("name", TypeCode.STRING), + ] + metadata = self._make_result_set_metadata(fields) + chunk1_rows = [[i, f"name_{i}"] for i in range(10)] + chunk2_rows = [[i, f"name_{i}"] for i in range(10, 20)] + values1 = [self._make_value(cell) for row in chunk1_rows for cell in row] + values2 = [self._make_value(cell) for row in chunk2_rows for cell in row] + + result_set1 = self._make_partial_result_set(values1, metadata=metadata) + result_set2 = self._make_partial_result_set(values2) + iterator = _MockCancellableIterator(result_set1, result_set2) + streamed = self._make_one(iterator) + stream_iter = iter(streamed) + first_part = [next(stream_iter) for _ in range(5)] + self.assertEqual(first_part, chunk1_rows[:5]) + middle_part = [next(stream_iter) for _ in range(10)] + self.assertEqual(middle_part, chunk1_rows[5:] + chunk2_rows[:5]) + final_part = list(stream_iter) + self.assertEqual(final_part, chunk2_rows[5:]) + + def test___iter___early_break(self): + from google.cloud.spanner_v1 import TypeCode + + fields = [ + self._make_scalar_field("id", TypeCode.INT64), + self._make_scalar_field("name", TypeCode.STRING), + ] + metadata = self._make_result_set_metadata(fields) + expected_rows = [[i, f"name_{i}"] for i in range(10)] + values = [self._make_value(cell) for row in expected_rows for cell in row] + + result_set = self._make_partial_result_set(values, metadata=metadata) + iterator = _MockCancellableIterator(result_set) + streamed = self._make_one(iterator) + consumed = [] + for row in streamed: + consumed.append(row) + if len(consumed) == 3: + break + + self.assertEqual(consumed, expected_rows[:3]) + + def test___iter___mid_stream_error(self): + from google.cloud.spanner_v1 import TypeCode + + fields = [ + self._make_scalar_field("id", TypeCode.INT64), + self._make_scalar_field("name", TypeCode.STRING), + ] + metadata = self._make_result_set_metadata(fields) + chunk1_rows = [[i, f"name_{i}"] for i in range(5)] + values1 = [self._make_value(cell) for row in chunk1_rows for cell in row] + result_set1 = self._make_partial_result_set(values1, metadata=metadata) + + def mock_iterator(): + yield result_set1 + raise RuntimeError("Stream error midway") + + streamed = self._make_one(mock_iterator()) + consumed = [] + with self.assertRaises(RuntimeError) as context: + for row in streamed: + consumed.append(row) + + self.assertEqual(consumed, chunk1_rows) + self.assertIn("Stream error midway", str(context.exception)) + class _MockCancellableIterator(object): cancel_calls = 0