fix(core): make abatch_iterate consistent with batch_iterate for None and zero size (#39367)

This commit is contained in:
pxmps authored and GitHub committed 2026-08-13 17:27:44 -04:00
1 parent 162f9d9582
commit f9ee55d94c
4 files changed
+50 -7

No files matched your search

+15 -5
View File
@@ -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
+6
View File
@@ -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))
+17 -1
View File
@@ -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]))]
+12 -1
View File
@@ -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]))