diff --git a/sqlite_utils/utils.py b/sqlite_utils/utils.py index ee6695b55..542ee290c 100644 --- a/sqlite_utils/utils.py +++ b/sqlite_utils/utils.py @@ -577,6 +577,8 @@ def chunks(sequence: Iterable[T], size: int) -> Iterable[Iterable[T]]: :param sequence: Any Python iterator :param size: The size of each chunk """ + if size < 1: + raise ValueError(f"chunk size must be at least 1, got {size}") iterator = iter(sequence) for item in iterator: yield itertools.chain([item], itertools.islice(iterator, size - 1)) diff --git a/tests/test_utils.py b/tests/test_utils.py index 360a4436a..2c608bddd 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -41,6 +41,16 @@ def test_chunks(size, expected): assert chunks == expected +def test_chunks_size_zero_raises(): + with pytest.raises(ValueError, match="chunk size must be at least 1"): + list(utils.chunks([1, 2, 3], 0)) + + +def test_chunks_negative_size_raises(): + with pytest.raises(ValueError, match="chunk size must be at least 1"): + list(utils.chunks([1, 2, 3], -1)) + + def test_hash_record(): expected = "d383e7c0ba88f5ffcdd09be660de164b3847401a" assert utils.hash_record({"name": "Cleo", "twitter": "CleoPaws"}) == expected