From a8f9cc6f64f299830834428509940d448b82b4ed Mon Sep 17 00:00:00 2001 From: Simon Willison Date: Sat, 8 Jan 2022 13:16:34 -0800 Subject: [PATCH] Add test for chunks(), refs #364 --- sqlite_utils/db.py | 7 +------ sqlite_utils/utils.py | 7 +++++++ tests/test_utils.py | 15 +++++++++++++++ 3 files changed, 23 insertions(+), 6 deletions(-) diff --git a/sqlite_utils/db.py b/sqlite_utils/db.py index 0ff056c..dfc4723 100644 --- a/sqlite_utils/db.py +++ b/sqlite_utils/db.py @@ -1,4 +1,5 @@ from .utils import ( + chunks, sqlite3, OperationalError, suggest_column_types, @@ -2995,12 +2996,6 @@ class View(Queryable): ) -def chunks(sequence, size): - iterator = iter(sequence) - for item in iterator: - yield itertools.chain([item], itertools.islice(iterator, size - 1)) - - def jsonify_if_needed(value): if isinstance(value, decimal.Decimal): return float(value) diff --git a/sqlite_utils/utils.py b/sqlite_utils/utils.py index caec976..3b4769a 100644 --- a/sqlite_utils/utils.py +++ b/sqlite_utils/utils.py @@ -3,6 +3,7 @@ import contextlib import csv import enum import io +import itertools import json import os from . import recipes @@ -306,3 +307,9 @@ def _compile_code(code, imports, variable="value"): globals[import_.split(".")[0]] = __import__(import_) exec(code_o, globals, locals) return locals["fn"] + + +def chunks(sequence, 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 a44b2e0..8630e28 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -25,3 +25,18 @@ def test_decode_base64_values(input, expected, should_be_is): def test_find_spatialite(): spatialite = utils.find_spatialite() assert spatialite is None or isinstance(spatialite, str) + + +@pytest.mark.parametrize( + "size,expected", + ( + (1, [["a"], ["b"], ["c"], ["d"]]), + (2, [["a", "b"], ["c", "d"]]), + (3, [["a", "b", "c"], ["d"]]), + (4, [["a", "b", "c", "d"]]), + ), +) +def test_chunks(size, expected): + input = ["a", "b", "c", "d"] + chunks = list(map(list, utils.chunks(input, size))) + assert chunks == expected