diff --git a/libs/core/langchain_core/utils/aiter.py b/libs/core/langchain_core/utils/aiter.py index 8d74a50fb7..245cad098b 100644 --- a/libs/core/langchain_core/utils/aiter.py +++ b/libs/core/langchain_core/utils/aiter.py @@ -323,25 +323,35 @@ class aclosing(AbstractAsyncContextManager[Any]): # noqa: N801 async def abatch_iterate( - size: int, iterable: AsyncIterable[T] + size: int | None, iterable: AsyncIterable[T] ) -> AsyncIterator[list[T]]: """Utility batching function for async iterables. Args: size: The size of the batch. + + If `None`, returns a single batch. iterable: The async iterable to batch. Yields: The batches. + + Raises: + ValueError: If `size` is not `None` and is not a positive integer. """ + if size is None: + single_batch = [el async for el in iterable] + if single_batch: + yield single_batch + return + if size <= 0: + msg = f"Batch size must be a positive integer, got {size}." + raise ValueError(msg) batch: list[T] = [] async for element in iterable: - if len(batch) < size: - batch.append(element) - + batch.append(element) if len(batch) >= size: yield batch batch = [] - if batch: yield batch diff --git a/libs/core/langchain_core/utils/iter.py b/libs/core/langchain_core/utils/iter.py index b24c5f213a..b89bddb9fc 100644 --- a/libs/core/langchain_core/utils/iter.py +++ b/libs/core/langchain_core/utils/iter.py @@ -214,7 +214,13 @@ def batch_iterate(size: int | None, iterable: Iterable[T]) -> Iterator[list[T]]: Yields: The batches of the iterable. + + Raises: + ValueError: If `size` is not `None` and is not a positive integer. """ + if size is not None and size <= 0: + msg = f"Batch size must be a positive integer, got {size}." + raise ValueError(msg) it = iter(iterable) while True: chunk = list(islice(it, size)) diff --git a/libs/core/tests/unit_tests/utils/test_aiter.py b/libs/core/tests/unit_tests/utils/test_aiter.py index 078fbdf8ac..a2f903a474 100644 --- a/libs/core/tests/unit_tests/utils/test_aiter.py +++ b/libs/core/tests/unit_tests/utils/test_aiter.py @@ -12,10 +12,14 @@ from langchain_core.utils.aiter import abatch_iterate (3, [10, 20, 30, 40, 50], [[10, 20, 30], [40, 50]]), (1, [100, 200, 300], [[100], [200], [300]]), (4, [], []), + (None, [1, 2, 3], [[1, 2, 3]]), + (None, [], []), ], ) async def test_abatch_iterate( - input_size: int, input_iterable: list[str], expected_output: list[list[str]] + input_size: int | None, + input_iterable: list[str], + expected_output: list[list[str]], ) -> None: """Test batching function.""" @@ -29,3 +33,15 @@ async def test_abatch_iterate( output = [el async for el in iterator_] assert output == expected_output + + +@pytest.mark.parametrize("input_size", [0, -1]) +async def test_abatch_iterate_invalid_size(input_size: int) -> None: + """Non-positive sizes should raise instead of silently discarding data.""" + + async def _to_async_iterable(iterable: list[int]) -> AsyncIterator[int]: + for item in iterable: + yield item + + with pytest.raises(ValueError, match="positive integer"): + _ = [el async for el in abatch_iterate(input_size, _to_async_iterable([1, 2]))] diff --git a/libs/core/tests/unit_tests/utils/test_iter.py b/libs/core/tests/unit_tests/utils/test_iter.py index 84e7882d48..a3494a3ae0 100644 --- a/libs/core/tests/unit_tests/utils/test_iter.py +++ b/libs/core/tests/unit_tests/utils/test_iter.py @@ -10,10 +10,21 @@ from langchain_core.utils.iter import batch_iterate (3, [10, 20, 30, 40, 50], [[10, 20, 30], [40, 50]]), (1, [100, 200, 300], [[100], [200], [300]]), (4, [], []), + (None, [1, 2, 3], [[1, 2, 3]]), + (None, [], []), ], ) def test_batch_iterate( - input_size: int, input_iterable: list[str], expected_output: list[list[str]] + input_size: int | None, + input_iterable: list[str], + expected_output: list[list[str]], ) -> None: """Test batching function.""" assert list(batch_iterate(input_size, input_iterable)) == expected_output + + +@pytest.mark.parametrize("input_size", [0, -1]) +def test_batch_iterate_invalid_size(input_size: int) -> None: + """Non-positive sizes should raise instead of silently discarding data.""" + with pytest.raises(ValueError, match="positive integer"): + list(batch_iterate(input_size, [1, 2, 3]))